diff --git a/.gitignore b/.gitignore index 603ccfcf..8dd24c0c 100644 --- a/.gitignore +++ b/.gitignore @@ -34,5 +34,6 @@ test_compact_storage/* test_working_memory/* *.code-workspace local_vector_store/* +chroma_vector_store/* bench_results/* meta_memory/* \ No newline at end of file diff --git a/bench/eval_reme.py b/bench/eval_reme_old.py similarity index 90% rename from bench/eval_reme.py rename to bench/eval_reme_old.py index 8e6c5e9e..51265da7 100644 --- a/bench/eval_reme.py +++ b/bench/eval_reme_old.py @@ -158,10 +158,11 @@ async def process_user_async( tmp_file = os.path.join(tmp_dir, f"{user_data['uuid']}.json") # Clear existing memories for this user - await reme.vector_store.delete_collection(f"reme_eval_{user_name}") + collection_name = f"reme_eval_{user_name}".replace(" ", "_").lower() + await reme.vector_store.delete_collection(collection_name) # Update collection name for this user - reme.vector_store.set_collection_name(f"reme_eval_{user_name}") + reme.vector_store.set_collection_name(collection_name) new_user_data = { "uuid": user_data["uuid"], @@ -177,35 +178,40 @@ async def process_user_async( # Add messages to ReMe dialogue = session["dialogue"] - # Parse timestamp and format as "YYYY-MM-DD HH:MM:SS" - date_format = "%b %d, %Y, %H:%M:%S" - # dt = datetime.strptime(session["start_time"], date_format).replace(tzinfo=timezone.utc) - # time_created = dt.strftime("%Y-%m-%d %H:%M:%S") - formatted_dialogue = [ { "role": turn["role"], "content": turn["content"], - "time_created": datetime.strptime(turn["timestamp"], date_format) + "time_created": datetime.strptime(turn["timestamp"], "%b %d, %Y, %H:%M:%S") .replace(tzinfo=timezone.utc) .strftime("%Y-%m-%d %H:%M:%S"), } for turn in dialogue ] - # Add memory - result, duration_ms = await add_memory_async( - reme=reme, - user_id=user_name, - messages=formatted_dialogue, - ) - memories = [] - for memory_modes in result: - for memory_mode in memory_modes: - if not isinstance(memory_mode, MemoryNode): - continue + # Add memory - process every 2 messages + result = [] + total_duration_ms = 0 + batch_size = 4 - memories.append(memory_mode.content) + for i in range(0, len(formatted_dialogue), batch_size): + batch = formatted_dialogue[i : i + batch_size] + batch_result, duration_ms = await add_memory_async( + reme=reme, + user_id=user_name, + messages=batch, + ) + result.extend(batch_result) + total_duration_ms += duration_ms + + duration_ms = total_duration_ms + + memories = [] + for memory_mode in result: + if not isinstance(memory_mode, MemoryNode): + continue + + memories.append(memory_mode.content) print(memories) diff --git a/bench/halumem/compute_stats_from_tmp.py b/bench/halumem/compute_stats_from_tmp.py new file mode 100644 index 00000000..3891c21a --- /dev/null +++ b/bench/halumem/compute_stats_from_tmp.py @@ -0,0 +1,556 @@ +""" +Compute statistics from existing tmp JSON files (Stage 1 results). + +This script assumes that process_user_stage1 has already been run and JSON files +are available in the tmp directory. It will: +1. Load all JSON files from tmp directory +2. Generate the combined JSONL file +3. Run Stage 2 evaluation (extraction only, no new API calls) +4. Aggregate results and compute metrics + +Usage: + python bench/halumem/compute_stats_from_tmp.py --tmp_dir bench_results/reme/tmp +""" + +import asyncio +import json +import os +import time +from datetime import datetime + +from loguru import logger + +from eval_tools import ( + evaluation_for_memory_accuracy, + evaluation_for_memory_integrity, + evaluation_for_question, + evaluation_for_update_memory, +) + + +def compute_f1(precision: float, recall: float) -> float: + """Compute F1-score from precision and recall.""" + if precision + recall == 0: + return 0.0 + return 2 * (precision * recall) / (precision + recall) + + +def iter_jsonl(file_path: str): + """Iterate over lines in a JSONL file.""" + with open(file_path, "r", encoding="utf-8") as f: + for line in f: + line = line.strip() + if line: + yield json.loads(line) + + +async def process_user_stage2(idx: int, user_data: dict): + """Stage 2: Extract evaluation results from user's sessions (already computed in Stage 1).""" + user_name = user_data["user_name"] + + eval_results = { + "memory_integrity_records": [], + "memory_accuracy_records": [], + "memory_update_records": [], + "question_answering_records": [], + } + + logger.info(f"[{idx}]{user_name}: Extracting evaluation results from sessions...") + + # Extract evaluation results from each session + for session in user_data["sessions"]: + if session.get("is_generated_qa_session", False): + continue + + if "evaluation_results" not in session: + logger.warning(f"[{idx}]{user_name}: Session missing evaluation_results, skipping...") + continue + + session_eval = session["evaluation_results"] + eval_results["memory_integrity_records"].extend(session_eval.get("memory_integrity_records", [])) + eval_results["memory_accuracy_records"].extend(session_eval.get("memory_accuracy_records", [])) + eval_results["memory_update_records"].extend(session_eval.get("memory_update_records", [])) + eval_results["question_answering_records"].extend(session_eval.get("question_answering_records", [])) + + logger.info( + f"[{idx}]{user_name}: Extracted {len(eval_results['memory_integrity_records'])} integrity, " + f"{len(eval_results['memory_accuracy_records'])} accuracy, " + f"{len(eval_results['memory_update_records'])} update, " + f"{len(eval_results['question_answering_records'])} QA records", + ) + + return eval_results + + +def aggregate_eval_results(eval_results): + """Aggregate evaluation results and compute metrics.""" + + # Memory Integrity Evaluation + memory_integrity_scores = 0 + memory_integrity_weighted_scores = 0 + memory_integrity_valid_num = 0 + memory_integrity_num = 0 + memory_integrity_weighted_valid_num = 0 + memory_integrity_weighted_num = 0 + interference_memory_scores = 0 + interference_memory_valid_num = 0 + interference_memory_num = 0 + + for item in eval_results["memory_integrity_records"]: + item["is_valid"] = True + + if item["memory_source"] != "interference": + memory_integrity_num += 1 + memory_integrity_weighted_num += item["importance"] + else: + interference_memory_num += 1 + + if item["memory_integrity_score"] is None: + item["is_valid"] = False + continue + + if item["memory_source"] != "interference": + if item["memory_integrity_score"] == 2: + memory_integrity_scores += 1 + memory_integrity_weighted_scores += 0.5 * item["memory_integrity_score"] * item["importance"] + memory_integrity_valid_num += 1 + memory_integrity_weighted_valid_num += item["importance"] + else: + if item["memory_integrity_score"] == 0: + interference_memory_scores += 1 + interference_memory_valid_num += 1 + + eval_results["overall_score"]["memory_integrity"]["recall(all)"] = ( + memory_integrity_scores / memory_integrity_num if memory_integrity_num > 0 else 0 + ) + eval_results["overall_score"]["memory_integrity"]["recall(valid)"] = ( + memory_integrity_scores / memory_integrity_valid_num if memory_integrity_valid_num > 0 else 0 + ) + eval_results["overall_score"]["memory_integrity"]["weighted_recall(all)"] = ( + memory_integrity_weighted_scores / memory_integrity_weighted_num if memory_integrity_weighted_num > 0 else 0 + ) + eval_results["overall_score"]["memory_integrity"]["weighted_recall(valid)"] = ( + memory_integrity_weighted_scores / memory_integrity_weighted_valid_num + if memory_integrity_weighted_valid_num > 0 + else 0 + ) + eval_results["overall_score"]["memory_integrity"][ + "memory_valid_importance_sum" + ] = memory_integrity_weighted_valid_num + eval_results["overall_score"]["memory_integrity"]["memory_importance_sum"] = memory_integrity_weighted_num + eval_results["overall_score"]["memory_integrity"]["memory_valid_num"] = memory_integrity_valid_num + eval_results["overall_score"]["memory_integrity"]["memory_num"] = memory_integrity_num + eval_results["overall_score"]["memory_accuracy"]["interference_accuracy(all)"] = ( + interference_memory_scores / interference_memory_num if interference_memory_num > 0 else 0 + ) + eval_results["overall_score"]["memory_accuracy"]["interference_accuracy(valid)"] = ( + interference_memory_scores / interference_memory_valid_num if interference_memory_valid_num > 0 else 0 + ) + eval_results["overall_score"]["memory_accuracy"]["interference_memory_valid_num"] = interference_memory_valid_num + eval_results["overall_score"]["memory_accuracy"]["interference_memory_num"] = interference_memory_num + + # Memory Accuracy Evaluation + target_memory_accuracy_scores = 0 + memory_accuracy_weighted_scores = 0 + target_memory_accuracy_valid_num = 0 + target_memory_accuracy_num = 0 + memory_accuracy_valid_num = 0 + memory_accuracy_num = 0 + + for item in eval_results["memory_accuracy_records"]: + item["is_valid"] = True + memory_accuracy_num += 1 + + if item["is_included_in_golden_memories"] in ["true", "True"]: + target_memory_accuracy_num += 1 + + if item["memory_accuracy_score"] is None: + item["is_valid"] = False + continue + + if item["is_included_in_golden_memories"] in ["true", "True"]: + target_memory_accuracy_scores += 0.5 * item["memory_accuracy_score"] + target_memory_accuracy_valid_num += 1 + + memory_accuracy_weighted_scores += 0.5 * item["memory_accuracy_score"] + memory_accuracy_valid_num += 1 + + eval_results["overall_score"]["memory_accuracy"]["target_accuracy(all)"] = ( + target_memory_accuracy_scores / target_memory_accuracy_num if target_memory_accuracy_num > 0 else 0 + ) + eval_results["overall_score"]["memory_accuracy"]["target_accuracy(valid)"] = ( + target_memory_accuracy_scores / target_memory_accuracy_valid_num if target_memory_accuracy_valid_num > 0 else 0 + ) + eval_results["overall_score"]["memory_accuracy"]["target_memory_valid_num"] = target_memory_accuracy_valid_num + eval_results["overall_score"]["memory_accuracy"]["target_memory_num"] = target_memory_accuracy_num + eval_results["overall_score"]["memory_accuracy"]["weighted_accuracy(all)"] = ( + memory_accuracy_weighted_scores / memory_accuracy_num if memory_accuracy_num > 0 else 0 + ) + eval_results["overall_score"]["memory_accuracy"]["weighted_accuracy(valid)"] = ( + memory_accuracy_weighted_scores / memory_accuracy_valid_num if memory_accuracy_valid_num > 0 else 0 + ) + eval_results["overall_score"]["memory_accuracy"]["memory_valid_num"] = memory_accuracy_valid_num + eval_results["overall_score"]["memory_accuracy"]["memory_num"] = memory_accuracy_num + + # Memory Extraction F1-score + eval_results["overall_score"]["memory_extraction_f1"] = compute_f1( + precision=eval_results["overall_score"]["memory_accuracy"]["target_accuracy(all)"], + recall=eval_results["overall_score"]["memory_integrity"]["recall(all)"], + ) + + # Memory Update Evaluation + correct_update_memory_num = 0 + hallucination_update_memory_num = 0 + omission_update_memory_num = 0 + other_update_memory_num = 0 + update_memory_num = 0 + update_memory_valid_num = 0 + + for item in eval_results["memory_update_records"]: + item["is_valid"] = True + update_memory_num += 1 + + if item["memory_update_type"] not in ["Correct", "Hallucination", "Omission", "Other"]: + item["is_valid"] = False + continue + + if item["memory_update_type"] == "Correct": + correct_update_memory_num += 1 + elif item["memory_update_type"] == "Hallucination": + hallucination_update_memory_num += 1 + elif item["memory_update_type"] == "Omission": + omission_update_memory_num += 1 + elif item["memory_update_type"] == "Other": + other_update_memory_num += 1 + + update_memory_valid_num += 1 + + if update_memory_num > 0: + eval_results["overall_score"]["memory_update"]["correct_update_memory_ratio(all)"] = ( + correct_update_memory_num / update_memory_num + ) + eval_results["overall_score"]["memory_update"]["hallucination_update_memory_ratio(all)"] = ( + hallucination_update_memory_num / update_memory_num + ) + eval_results["overall_score"]["memory_update"]["omission_update_memory_ratio(all)"] = ( + omission_update_memory_num / update_memory_num + ) + eval_results["overall_score"]["memory_update"]["other_update_memory_ratio(all)"] = ( + other_update_memory_num / update_memory_num + ) + else: + eval_results["overall_score"]["memory_update"]["correct_update_memory_ratio(all)"] = 0 + eval_results["overall_score"]["memory_update"]["hallucination_update_memory_ratio(all)"] = 0 + eval_results["overall_score"]["memory_update"]["omission_update_memory_ratio(all)"] = 0 + eval_results["overall_score"]["memory_update"]["other_update_memory_ratio(all)"] = 0 + + if update_memory_valid_num > 0: + eval_results["overall_score"]["memory_update"]["correct_update_memory_ratio(valid)"] = ( + correct_update_memory_num / update_memory_valid_num + ) + eval_results["overall_score"]["memory_update"]["hallucination_update_memory_ratio(valid)"] = ( + hallucination_update_memory_num / update_memory_valid_num + ) + eval_results["overall_score"]["memory_update"]["omission_update_memory_ratio(valid)"] = ( + omission_update_memory_num / update_memory_valid_num + ) + eval_results["overall_score"]["memory_update"]["other_update_memory_ratio(valid)"] = ( + other_update_memory_num / update_memory_valid_num + ) + else: + eval_results["overall_score"]["memory_update"]["correct_update_memory_ratio(valid)"] = 0 + eval_results["overall_score"]["memory_update"]["hallucination_update_memory_ratio(valid)"] = 0 + eval_results["overall_score"]["memory_update"]["omission_update_memory_ratio(valid)"] = 0 + eval_results["overall_score"]["memory_update"]["other_update_memory_ratio(valid)"] = 0 + + eval_results["overall_score"]["memory_update"]["update_memory_valid_num"] = update_memory_valid_num + eval_results["overall_score"]["memory_update"]["update_memory_num"] = update_memory_num + + # Question-Answering Evaluation + correct_qa_num = 0 + hallucination_qa_num = 0 + omission_qa_num = 0 + qa_num = 0 + qa_valid_num = 0 + + for item in eval_results["question_answering_records"]: + item["is_valid"] = True + qa_num += 1 + + if item["result_type"] not in ["Correct", "Hallucination", "Omission"]: + item["is_valid"] = False + continue + + if item["result_type"] == "Correct": + correct_qa_num += 1 + elif item["result_type"] == "Hallucination": + hallucination_qa_num += 1 + elif item["result_type"] == "Omission": + omission_qa_num += 1 + + qa_valid_num += 1 + + if qa_num > 0: + eval_results["overall_score"]["question_answering"]["correct_qa_ratio(all)"] = correct_qa_num / qa_num + eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(all)"] = ( + hallucination_qa_num / qa_num + ) + eval_results["overall_score"]["question_answering"]["omission_qa_ratio(all)"] = omission_qa_num / qa_num + else: + eval_results["overall_score"]["question_answering"]["correct_qa_ratio(all)"] = 0 + eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(all)"] = 0 + eval_results["overall_score"]["question_answering"]["omission_qa_ratio(all)"] = 0 + + if qa_valid_num > 0: + eval_results["overall_score"]["question_answering"]["correct_qa_ratio(valid)"] = correct_qa_num / qa_valid_num + eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(valid)"] = ( + hallucination_qa_num / qa_valid_num + ) + eval_results["overall_score"]["question_answering"]["omission_qa_ratio(valid)"] = omission_qa_num / qa_valid_num + else: + eval_results["overall_score"]["question_answering"]["correct_qa_ratio(valid)"] = 0 + eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(valid)"] = 0 + eval_results["overall_score"]["question_answering"]["omission_qa_ratio(valid)"] = 0 + + eval_results["overall_score"]["question_answering"]["qa_valid_num"] = qa_valid_num + eval_results["overall_score"]["question_answering"]["qa_num"] = qa_num + + # Memory Type Accuracy + for item in eval_results["memory_integrity_records"]: + if "memory_integrity_score" not in item or "importance" not in item: + continue + score = 1 if item["memory_integrity_score"] == 2 else 0 + eval_results["overall_score"]["memory_type_accuracy"][item["memory_type"]]["memory_integrity_acc"] += score + eval_results["overall_score"]["memory_type_accuracy"][item["memory_type"]]["total_num"] += 1 + + for item in eval_results["memory_update_records"]: + if "memory_update_type" not in item or "importance" not in item: + continue + score = 1 if item["memory_update_type"] == "Correct" else 0 + eval_results["overall_score"]["memory_type_accuracy"][item["memory_type"]]["memory_update_acc"] += score + eval_results["overall_score"]["memory_type_accuracy"][item["memory_type"]]["total_num"] += 1 + + for key in eval_results["overall_score"]["memory_type_accuracy"]: + if eval_results["overall_score"]["memory_type_accuracy"][key]["total_num"] > 0: + total = eval_results["overall_score"]["memory_type_accuracy"][key]["total_num"] + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_integrity_acc"] = ( + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_integrity_acc"] / total + ) + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_update_acc"] = ( + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_update_acc"] / total + ) + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_acc"] = ( + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_integrity_acc"] + + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_update_acc"] + ) + else: + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_integrity_acc"] = 0 + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_update_acc"] = 0 + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_acc"] = 0 + + return eval_results + + +async def main_async(tmp_dir: str): + """Main function to compute statistics from tmp directory.""" + start_time = time.time() + + # Determine paths + parent_dir = os.path.dirname(tmp_dir) + frame = "reme" + + output_file_stage1 = os.path.join(parent_dir, f"{frame}_eval_results.jsonl") + output_file_stage2 = os.path.join(parent_dir, f"{frame}_eval_stat_result.json") + + print("\n" + "=" * 80) + print("LOADING STAGE 1 RESULTS FROM TMP DIRECTORY") + print(f"Tmp Directory: {tmp_dir}") + print("=" * 80) + + # Step 1: Combine all tmp JSON files into the stage1 output JSONL + json_files = [f for f in os.listdir(tmp_dir) if f.endswith(".json")] + print(f"\n📁 Found {len(json_files)} JSON files in tmp directory") + + with open(output_file_stage1, "w", encoding="utf-8") as f_out: + for file_name in json_files: + file_path = os.path.join(tmp_dir, file_name) + with open(file_path, "r", encoding="utf-8") as f_in: + data = json.load(f_in) + f_out.write(json.dumps(data, ensure_ascii=False) + "\n") + + print(f"✅ Combined results saved to: {output_file_stage1}") + + # Step 2: Run Stage 2 evaluation (extraction only) + print("\n" + "=" * 80) + print("STAGE 2: EXTRACTING AND AGGREGATING EVALUATION RESULTS") + print("=" * 80) + + tmp_dir2 = os.path.join(parent_dir, "tmp2") + os.makedirs(tmp_dir2, exist_ok=True) + + start_stage2 = time.time() + + # Load all users and process + user_data_list = list(enumerate(iter_jsonl(output_file_stage1), 1)) + + for idx, user_data in user_data_list: + uuid = user_data["uuid"] + tmp_file = os.path.join(tmp_dir2, f"{uuid}.json") + + if os.path.exists(tmp_file): + print(f"⚡ Skipping user {uuid} ({idx}/{len(user_data_list)}) — cached result found.") + continue + + print(f"[{idx}/{len(user_data_list)}] Processing user {uuid}...") + t_user_result = await process_user_stage2(idx, user_data) + + with open(tmp_file, "w", encoding="utf-8") as f: + json.dump(t_user_result, f, ensure_ascii=False, indent=4) + + elapsed = time.time() - start_stage2 + print(f"[{idx}/{len(user_data_list)}] ✅ Finished user {uuid}, elapsed {elapsed:.2f}s.") + + # Calculate time consuming + add_dialogue_duration_time = 0 + search_memory_duration_time = 0 + + for user_data in iter_jsonl(output_file_stage1): + sessions = user_data["sessions"] + + for session in sessions: + if "add_dialogue_duration_ms" in session: + add_dialogue_duration_time += session["add_dialogue_duration_ms"] + + if "questions" in session: + for question in session["questions"]: + if "search_duration_ms" in question: + search_memory_duration_time += question["search_duration_ms"] + + add_dialogue_duration_time = add_dialogue_duration_time / 1000 / 60 + search_memory_duration_time = search_memory_duration_time / 1000 / 60 + + print("\n🔄 Aggregating all user results...") + + eval_results = { + "overall_score": { + "memory_integrity": {}, + "memory_accuracy": {}, + "memory_extraction_f1": 0, + "memory_update": {}, + "question_answering": {}, + "memory_type_accuracy": { + "Event Memory": { + "memory_integrity_acc": 0, + "memory_update_acc": 0, + "total_num": 0, + }, + "Persona Memory": { + "memory_integrity_acc": 0, + "memory_update_acc": 0, + "total_num": 0, + }, + "Relationship Memory": { + "memory_integrity_acc": 0, + "memory_update_acc": 0, + "total_num": 0, + }, + }, + "time_consuming": { + "add_dialogue_duration_time": add_dialogue_duration_time, + "search_memory_duration_time": search_memory_duration_time, + "total_duration_time": add_dialogue_duration_time + search_memory_duration_time, + }, + }, + "memory_integrity_records": [], + "memory_accuracy_records": [], + "memory_update_records": [], + "question_answering_records": [], + } + + for file_name in os.listdir(tmp_dir2): + if not file_name.endswith(".json"): + continue + user_file = os.path.join(tmp_dir2, file_name) + with open(user_file, "r", encoding="utf-8") as f: + user_result = json.load(f) + + eval_results["memory_accuracy_records"].extend(user_result.get("memory_accuracy_records", [])) + eval_results["memory_integrity_records"].extend(user_result.get("memory_integrity_records", [])) + eval_results["memory_update_records"].extend(user_result.get("memory_update_records", [])) + eval_results["question_answering_records"].extend(user_result.get("question_answering_records", [])) + + eval_results = aggregate_eval_results(eval_results) + + with open(output_file_stage2, "w", encoding="utf-8") as f: + json.dump(eval_results, f, ensure_ascii=False, indent=4) + + elapsed_total = time.time() - start_time + print(f"\n✅ All done in {elapsed_total:.2f}s. Results saved to {output_file_stage2}") + + # Print summary + print("\n" + "=" * 80) + print("EVALUATION SUMMARY") + print("=" * 80) + print("\n📊 Memory Integrity:") + print(f" - Recall (all): {eval_results['overall_score']['memory_integrity'].get('recall(all)', 0):.4f}") + print(f" - Recall (valid): {eval_results['overall_score']['memory_integrity'].get('recall(valid)', 0):.4f}") + print(f" - Weighted Recall (all): " + f"{eval_results['overall_score']['memory_integrity'].get('weighted_recall(all)', 0):.4f}") + + print(f"\n📊 Memory Accuracy:") + print(f" - Target Accuracy (all): {eval_results['overall_score']['memory_accuracy'].get('target_accuracy(all)', 0):.4f}") + print(f" - Target Accuracy (valid): {eval_results['overall_score']['memory_accuracy'].get('target_accuracy(valid)', 0):.4f}", + ) + print( + f" - Weighted Accuracy (all): {eval_results['overall_score']['memory_accuracy'].get('weighted_accuracy(all)', 0):.4f}", + ) + + print(f"\n📊 Memory Extraction F1: {eval_results['overall_score']['memory_extraction_f1']:.4f}") + + print(f"\n📊 Memory Update:") + print( + f" - Correct (all): {eval_results['overall_score']['memory_update'].get('correct_update_memory_ratio(all)', 0):.4f}", + ) + print( + f" - Hallucination (all): {eval_results['overall_score']['memory_update'].get('hallucination_update_memory_ratio(all)', 0):.4f}", + ) + print( + f" - Omission (all): {eval_results['overall_score']['memory_update'].get('omission_update_memory_ratio(all)', 0):.4f}", + ) + + print(f"\n📊 Question Answering:") + print( + f" - Correct (all): {eval_results['overall_score']['question_answering'].get('correct_qa_ratio(all)', 0):.4f}", + ) + print( + f" - Hallucination (all): {eval_results['overall_score']['question_answering'].get('hallucination_qa_ratio(all)', 0):.4f}", + ) + print( + f" - Omission (all): {eval_results['overall_score']['question_answering'].get('omission_qa_ratio(all)', 0):.4f}", + ) + + print(f"\n⏱️ Time Consuming:") + print(f" - Add Dialogue: {add_dialogue_duration_time:.2f} min") + print(f" - Search Memory: {search_memory_duration_time:.2f} min") + print(f" - Total: {add_dialogue_duration_time + search_memory_duration_time:.2f} min") + print("=" * 80) + + +def main(tmp_dir: str): + """Synchronous entry point.""" + asyncio.run(main_async(tmp_dir)) + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser(description="Compute statistics from existing Stage 1 tmp results") + parser.add_argument( + "--tmp_dir", + type=str, + required=True, + help="Path to tmp directory containing Stage 1 JSON results (e.g., bench_results/reme/tmp)", + ) + args = parser.parse_args() + + main(tmp_dir=args.tmp_dir) diff --git a/bench/halumem/eval_reme.py b/bench/halumem/eval_reme.py new file mode 100644 index 00000000..d3044535 --- /dev/null +++ b/bench/halumem/eval_reme.py @@ -0,0 +1,925 @@ +""" +Complete evaluation script for ReMe on HaluMem benchmark. + +This script performs the full evaluation pipeline: +1. Load HaluMem data +2. Process each user's sessions with ReMe (summary + retrieve) - Stage 1 (Parallel) +3. Evaluate memory integrity, accuracy, updates, and question answering - Stage 2 (Sequential) +4. Generate metrics and statistics + +Usage: + python bench/halumem/eval_reme.py --data_path /Users/yuli/workspace/HaluMem/data/HaluMem-Long.jsonl \ + --top_k 20 --user_num 100 --max_concurrency 20 + python bench/halumem/eval_reme.py --data_path ./HaluMem-Long.jsonl \ + --top_k 20 --user_num 100 --max_concurrency 20 + + python bench/halumem/eval_reme.py --data_path /Users/yuli/workspace/HaluMem/data/tmp_14.jsonl \ + --top_k 20 --user_num 1 --max_concurrency 1 +""" + +import asyncio +import copy +import json +import os +import re +import time +from datetime import datetime, timezone + +from loguru import logger + +from eval_tools import ( + _PROMPTS, + evaluation_for_memory_accuracy, + evaluation_for_memory_integrity, + evaluation_for_question, + evaluation_for_update_memory, +) +from llms import llm_request +from reme_ai.core.enumeration import MemoryType +from reme_ai.core.schema import MemoryNode +from reme_ai.reme import ReMe + +# Template for formatting memories (from shared YAML config) +TEMPLATE_MEMOS = _PROMPTS["TEMPLATE_MEMOS"] + +# Prompt for question answering (using optimized PROMPT_MEMOS) +PROMPT_MEMOS = _PROMPTS["PROMPT_MEMOS"] + +# Initialize ReMe with rate limiting configuration +# The default LLM can be overridden at call time using model_name parameter +reme: ReMe = ReMe() + + +def extract_user_name(persona_info: str): + """Extract user name from persona info.""" + match = re.search(r"Name:\s*(.*?); Gender:", persona_info) + if match: + username = match.group(1).strip() + return username + else: + raise ValueError("No name found.") + + +def iter_jsonl(file_path: str): + """Iterate over lines in a JSONL file.""" + with open(file_path, "r", encoding="utf-8") as f: + for line in f: + line = line.strip() + if line: + yield json.loads(line) + + +def compute_f1(precision: float, recall: float) -> float: + """Compute F1-score from precision and recall.""" + if precision + recall == 0: + return 0.0 + return 2 * (precision * recall) / (precision + recall) + + +# ==================== Stage 1: Data Processing ==================== + + +async def add_memory_async(user_id: str, messages: list[dict]) -> tuple[list[MemoryNode], float]: + """Add memory to ReMe system asynchronously.""" + start = time.time() + result = await reme.summary_v2(messages=messages, user_id=user_id) + duration_ms = (time.time() - start) * 1000 + return result, duration_ms + + +async def search_memory_async(query: str, user_id: str, top_k: int = 20): + """Search memory from ReMe system asynchronously.""" + start = time.time() + memories = await reme.retrieve_v2(query=query, user_id=user_id, top_k=top_k) + + # Format the context + context = TEMPLATE_MEMOS.format(user_id=user_id, memories=memories) + duration_ms = (time.time() - start) * 1000 + return context, memories, duration_ms + + +async def process_user_stage1( + user_data: dict, + top_k_value: int, + save_path: str, +): + """Stage 1: Process user data through ReMe (summary + retrieve).""" + user_name = extract_user_name(user_data["persona_info"]) + sessions = user_data["sessions"] + + tmp_dir = os.path.join(save_path, "tmp") + os.makedirs(tmp_dir, exist_ok=True) + tmp_file = os.path.join(tmp_dir, f"{user_data['uuid']}.json") + + new_user_data = { + "uuid": user_data["uuid"], + "user_name": user_name, + "sessions": [], + } + + for idx, session in enumerate(sessions): + logger.info(f"Processing user {user_name}: session {idx}/{len(sessions)}") + new_session = { + "memory_points": session["memory_points"], + "dialogue": session["dialogue"], + } + + # Format dialogue + dialogue = session["dialogue"] + formatted_dialogue = [ + { + "role": turn["role"], + "content": turn["content"], + "time_created": datetime.strptime(turn["timestamp"], "%b %d, %Y, %H:%M:%S") + .replace(tzinfo=timezone.utc) + .strftime("%Y-%m-%d %H:%M:%S"), + } + for turn in dialogue + ] + + # Process in batches + result = [] + total_duration_ms = 0 + batch_size = 20 + + for i in range(0, len(formatted_dialogue), batch_size): + batch = formatted_dialogue[i : i + batch_size] + batch_result, duration_ms = await add_memory_async( + user_id=user_name, + messages=batch, + ) + if batch_result: + result.extend(batch_result) + total_duration_ms += duration_ms + + duration_ms = total_duration_ms + + # Extract memory content + memories = [] + for memory_node in result: + if isinstance(memory_node, MemoryNode) and memory_node.memory_type is not MemoryType.HISTORY: + memories.append(memory_node.content) + + if session.get("is_generated_qa_session", False): + new_session["add_dialogue_duration_ms"] = duration_ms + new_session["is_generated_qa_session"] = True + del new_session["dialogue"] + del new_session["memory_points"] + new_user_data["sessions"].append(new_session) + continue + + # Store extracted memories + new_session["extracted_memories"] = memories + new_session["add_dialogue_duration_ms"] = duration_ms + + # Search updated memories for memory points + for memory in new_session["memory_points"]: + if memory["is_update"] == "False" or not memory.get("original_memories"): + continue + + _, memories_from_system, duration_ms = await search_memory_async( + query=memory["memory_content"], + user_id=user_name, + top_k=10, + ) + + memory["memories_from_system"] = memories_from_system + + # Process questions + if "questions" not in session: + new_user_data["sessions"].append(new_session) + continue + + new_session["questions"] = [] + + for qa in session["questions"]: + context, _, duration_ms = await search_memory_async( + query=qa["question"], + user_id=user_name, + top_k=top_k_value, + ) + + new_qa = copy.deepcopy(qa) + new_qa["context"] = context + new_qa["search_duration_ms"] = duration_ms + + prompt = PROMPT_MEMOS.format( + context=context, + question=qa["question"], + ) + + start_time = time.time() + response = await llm_request(prompt) + new_qa["system_response"] = response + new_qa["response_duration_ms"] = (time.time() - start_time) * 1000 + + new_session["questions"].append(new_qa) + + # ==================== Evaluation for this session ==================== + session_eval_results = { + "memory_integrity_records": [], + "memory_accuracy_records": [], + "memory_update_records": [], + "question_answering_records": [], + } + + uuid = user_data["uuid"] + golden_memories = session["memory_points"] + extract_memories = new_session["extracted_memories"] + extract_memories_str = "\n".join(extract_memories) + + # Evaluate Memory Integrity + logger.info(f"Evaluating Memory Integrity for session {idx}...") + for memory in golden_memories: + if memory["is_update"] == "True" and memory.get("memories_from_system", []): + # Skip update memories for integrity check + continue + + new_memory = copy.deepcopy(memory) + new_memory["uuid"] = uuid + new_memory["session_id"] = idx + + if extract_memories_str.strip() == "": + new_memory["memory_integrity_score"] = 0 + new_memory["memory_integrity_reasoning"] = "No memories extracted" + session_eval_results["memory_integrity_records"].append(new_memory) + continue + + result = await evaluation_for_memory_integrity(extract_memories_str, memory["memory_content"]) + score = int(result.get("score")) + reasoning = result.get("reasoning", "") + new_memory["memory_integrity_score"] = score + new_memory["memory_integrity_reasoning"] = reasoning + session_eval_results["memory_integrity_records"].append(new_memory) + + # Evaluate Memory Accuracy + logger.info(f"Evaluating Memory Accuracy for session {idx}...") + dialogue = session["dialogue"] + dialogue_str = [] + for turn in dialogue: + dialogue_str.append(f'[{turn["timestamp"]}]{turn["role"]}: {turn["content"]}') + if turn["role"] == "assistant": + dialogue_str.append("") + dialogue_str = "\n".join(dialogue_str) + + golden_memories_str = "\n".join( + [m["memory_content"] for m in golden_memories if m["memory_source"] != "interference"], + ) + + for memory in extract_memories: + new_memory = { + "uuid": uuid, + "session_id": idx, + "memory_content": memory, + } + result = await evaluation_for_memory_accuracy(dialogue_str, golden_memories_str, memory) + score = int(result.get("accuracy_score")) + is_included_in_golden_memories = result.get("is_included_in_golden_memories", "false") + reason = result.get("reason", "") + new_memory["memory_accuracy_score"] = score + new_memory["is_included_in_golden_memories"] = is_included_in_golden_memories + new_memory["memory_accuracy_reason"] = reason + session_eval_results["memory_accuracy_records"].append(new_memory) + + # Evaluate Memory Update + logger.info(f"Evaluating Memory Update for session {idx}...") + for memory in golden_memories: + if memory["is_update"] == "False" or not memory.get("original_memories"): + continue + + if not memory.get("memories_from_system", []): + continue + + update_memory = copy.deepcopy(memory) + update_memory["uuid"] = uuid + update_memory["session_id"] = idx + + result = await evaluation_for_update_memory( + "\n".join(update_memory["memories_from_system"]), + update_memory["memory_content"], + "\n".join(update_memory["original_memories"]), + ) + update_type = result.get("evaluation_result") + reason = result.get("reason", "") + update_memory["memory_update_type"] = update_type + update_memory["memory_update_reason"] = reason + session_eval_results["memory_update_records"].append(update_memory) + + # Evaluate Question Answering + if "questions" in new_session: + logger.info(f"Evaluating Question Answering for session {idx}...") + for qa in new_session["questions"]: + new_qa = copy.deepcopy(qa) + new_qa["uuid"] = uuid + new_qa["session_id"] = idx + + result = await evaluation_for_question( + qa["question"], + qa["answer"], + "\n".join([i["memory_content"] for i in qa["evidence"]]), + qa["system_response"], + ) + result_type = result.get("evaluation_result") + reasoning = result.get("reasoning", "") + new_qa["result_type"] = result_type + new_qa["question_answering_reasoning"] = reasoning + session_eval_results["question_answering_records"].append(new_qa) + + # Store evaluation results in session + new_session["evaluation_results"] = session_eval_results + + new_user_data["sessions"].append(new_session) + + # Save results + with open(tmp_file, "w", encoding="utf-8") as f: + json.dump(new_user_data, f, ensure_ascii=False, indent=2) + session_size = len(new_user_data["sessions"]) + logger.info(f"✅ Saved user {user_name} to {tmp_file} session_size={session_size}") + + logger.info(f"✅ Saved user {user_name} to {tmp_file} all!") + return {"uuid": user_data["uuid"], "status": "ok", "path": tmp_file} + + +# ==================== Stage 2: Evaluation ==================== + + +async def process_user_stage2(idx: int, user_data: dict): + """Stage 2: Extract evaluation results from user's sessions (already computed in Stage 1).""" + user_name = user_data["user_name"] + + eval_results = { + "memory_integrity_records": [], + "memory_accuracy_records": [], + "memory_update_records": [], + "question_answering_records": [], + } + + logger.info(f"[{idx}]{user_name}: Extracting evaluation results from sessions...") + + # Extract evaluation results from each session + for session in user_data["sessions"]: + if session.get("is_generated_qa_session", False): + continue + + if "evaluation_results" not in session: + logger.warning(f"[{idx}]{user_name}: Session missing evaluation_results, skipping...") + continue + + session_eval = session["evaluation_results"] + eval_results["memory_integrity_records"].extend(session_eval.get("memory_integrity_records", [])) + eval_results["memory_accuracy_records"].extend(session_eval.get("memory_accuracy_records", [])) + eval_results["memory_update_records"].extend(session_eval.get("memory_update_records", [])) + eval_results["question_answering_records"].extend(session_eval.get("question_answering_records", [])) + + logger.info( + f"[{idx}]{user_name}: Extracted {len(eval_results['memory_integrity_records'])} integrity, " + f"{len(eval_results['memory_accuracy_records'])} accuracy, " + f"{len(eval_results['memory_update_records'])} update, " + f"{len(eval_results['question_answering_records'])} QA records", + ) + + return eval_results + + +def aggregate_eval_results(eval_results): + """Aggregate evaluation results and compute metrics.""" + + # Memory Integrity Evaluation + memory_integrity_scores = 0 + memory_integrity_weighted_scores = 0 + memory_integrity_valid_num = 0 + memory_integrity_num = 0 + memory_integrity_weighted_valid_num = 0 + memory_integrity_weighted_num = 0 + interference_memory_scores = 0 + interference_memory_valid_num = 0 + interference_memory_num = 0 + + for item in eval_results["memory_integrity_records"]: + item["is_valid"] = True + + if item["memory_source"] != "interference": + memory_integrity_num += 1 + memory_integrity_weighted_num += item["importance"] + else: + interference_memory_num += 1 + + if item["memory_integrity_score"] is None: + item["is_valid"] = False + continue + + if item["memory_source"] != "interference": + if item["memory_integrity_score"] == 2: + memory_integrity_scores += 1 + memory_integrity_weighted_scores += 0.5 * item["memory_integrity_score"] * item["importance"] + memory_integrity_valid_num += 1 + memory_integrity_weighted_valid_num += item["importance"] + else: + if item["memory_integrity_score"] == 0: + interference_memory_scores += 1 + interference_memory_valid_num += 1 + + eval_results["overall_score"]["memory_integrity"]["recall(all)"] = ( + memory_integrity_scores / memory_integrity_num if memory_integrity_num > 0 else 0 + ) + eval_results["overall_score"]["memory_integrity"]["recall(valid)"] = ( + memory_integrity_scores / memory_integrity_valid_num if memory_integrity_valid_num > 0 else 0 + ) + eval_results["overall_score"]["memory_integrity"]["weighted_recall(all)"] = ( + memory_integrity_weighted_scores / memory_integrity_weighted_num if memory_integrity_weighted_num > 0 else 0 + ) + eval_results["overall_score"]["memory_integrity"]["weighted_recall(valid)"] = ( + memory_integrity_weighted_scores / memory_integrity_weighted_valid_num + if memory_integrity_weighted_valid_num > 0 + else 0 + ) + eval_results["overall_score"]["memory_integrity"][ + "memory_valid_importance_sum" + ] = memory_integrity_weighted_valid_num + eval_results["overall_score"]["memory_integrity"]["memory_importance_sum"] = memory_integrity_weighted_num + eval_results["overall_score"]["memory_integrity"]["memory_valid_num"] = memory_integrity_valid_num + eval_results["overall_score"]["memory_integrity"]["memory_num"] = memory_integrity_num + eval_results["overall_score"]["memory_accuracy"]["interference_accuracy(all)"] = ( + interference_memory_scores / interference_memory_num if interference_memory_num > 0 else 0 + ) + eval_results["overall_score"]["memory_accuracy"]["interference_accuracy(valid)"] = ( + interference_memory_scores / interference_memory_valid_num if interference_memory_valid_num > 0 else 0 + ) + eval_results["overall_score"]["memory_accuracy"]["interference_memory_valid_num"] = interference_memory_valid_num + eval_results["overall_score"]["memory_accuracy"]["interference_memory_num"] = interference_memory_num + + # Memory Accuracy Evaluation + target_memory_accuracy_scores = 0 + memory_accuracy_weighted_scores = 0 + target_memory_accuracy_valid_num = 0 + target_memory_accuracy_num = 0 + memory_accuracy_valid_num = 0 + memory_accuracy_num = 0 + + for item in eval_results["memory_accuracy_records"]: + item["is_valid"] = True + memory_accuracy_num += 1 + + if item["is_included_in_golden_memories"] in ["true", "True"]: + target_memory_accuracy_num += 1 + + if item["memory_accuracy_score"] is None: + item["is_valid"] = False + continue + + if item["is_included_in_golden_memories"] in ["true", "True"]: + target_memory_accuracy_scores += 0.5 * item["memory_accuracy_score"] + target_memory_accuracy_valid_num += 1 + + memory_accuracy_weighted_scores += 0.5 * item["memory_accuracy_score"] + memory_accuracy_valid_num += 1 + + eval_results["overall_score"]["memory_accuracy"]["target_accuracy(all)"] = ( + target_memory_accuracy_scores / target_memory_accuracy_num if target_memory_accuracy_num > 0 else 0 + ) + eval_results["overall_score"]["memory_accuracy"]["target_accuracy(valid)"] = ( + target_memory_accuracy_scores / target_memory_accuracy_valid_num if target_memory_accuracy_valid_num > 0 else 0 + ) + eval_results["overall_score"]["memory_accuracy"]["target_memory_valid_num"] = target_memory_accuracy_valid_num + eval_results["overall_score"]["memory_accuracy"]["target_memory_num"] = target_memory_accuracy_num + eval_results["overall_score"]["memory_accuracy"]["weighted_accuracy(all)"] = ( + memory_accuracy_weighted_scores / memory_accuracy_num if memory_accuracy_num > 0 else 0 + ) + eval_results["overall_score"]["memory_accuracy"]["weighted_accuracy(valid)"] = ( + memory_accuracy_weighted_scores / memory_accuracy_valid_num if memory_accuracy_valid_num > 0 else 0 + ) + eval_results["overall_score"]["memory_accuracy"]["memory_valid_num"] = memory_accuracy_valid_num + eval_results["overall_score"]["memory_accuracy"]["memory_num"] = memory_accuracy_num + + # Memory Extraction F1-score + eval_results["overall_score"]["memory_extraction_f1"] = compute_f1( + precision=eval_results["overall_score"]["memory_accuracy"]["target_accuracy(all)"], + recall=eval_results["overall_score"]["memory_integrity"]["recall(all)"], + ) + + # Memory Update Evaluation + correct_update_memory_num = 0 + hallucination_update_memory_num = 0 + omission_update_memory_num = 0 + other_update_memory_num = 0 + update_memory_num = 0 + update_memory_valid_num = 0 + + for item in eval_results["memory_update_records"]: + item["is_valid"] = True + update_memory_num += 1 + + if item["memory_update_type"] not in ["Correct", "Hallucination", "Omission", "Other"]: + item["is_valid"] = False + continue + + if item["memory_update_type"] == "Correct": + correct_update_memory_num += 1 + elif item["memory_update_type"] == "Hallucination": + hallucination_update_memory_num += 1 + elif item["memory_update_type"] == "Omission": + omission_update_memory_num += 1 + elif item["memory_update_type"] == "Other": + other_update_memory_num += 1 + + update_memory_valid_num += 1 + + if update_memory_num > 0: + eval_results["overall_score"]["memory_update"]["correct_update_memory_ratio(all)"] = ( + correct_update_memory_num / update_memory_num + ) + eval_results["overall_score"]["memory_update"]["hallucination_update_memory_ratio(all)"] = ( + hallucination_update_memory_num / update_memory_num + ) + eval_results["overall_score"]["memory_update"]["omission_update_memory_ratio(all)"] = ( + omission_update_memory_num / update_memory_num + ) + eval_results["overall_score"]["memory_update"]["other_update_memory_ratio(all)"] = ( + other_update_memory_num / update_memory_num + ) + else: + eval_results["overall_score"]["memory_update"]["correct_update_memory_ratio(all)"] = 0 + eval_results["overall_score"]["memory_update"]["hallucination_update_memory_ratio(all)"] = 0 + eval_results["overall_score"]["memory_update"]["omission_update_memory_ratio(all)"] = 0 + eval_results["overall_score"]["memory_update"]["other_update_memory_ratio(all)"] = 0 + + if update_memory_valid_num > 0: + eval_results["overall_score"]["memory_update"]["correct_update_memory_ratio(valid)"] = ( + correct_update_memory_num / update_memory_valid_num + ) + eval_results["overall_score"]["memory_update"]["hallucination_update_memory_ratio(valid)"] = ( + hallucination_update_memory_num / update_memory_valid_num + ) + eval_results["overall_score"]["memory_update"]["omission_update_memory_ratio(valid)"] = ( + omission_update_memory_num / update_memory_valid_num + ) + eval_results["overall_score"]["memory_update"]["other_update_memory_ratio(valid)"] = ( + other_update_memory_num / update_memory_valid_num + ) + else: + eval_results["overall_score"]["memory_update"]["correct_update_memory_ratio(valid)"] = 0 + eval_results["overall_score"]["memory_update"]["hallucination_update_memory_ratio(valid)"] = 0 + eval_results["overall_score"]["memory_update"]["omission_update_memory_ratio(valid)"] = 0 + eval_results["overall_score"]["memory_update"]["other_update_memory_ratio(valid)"] = 0 + + eval_results["overall_score"]["memory_update"]["update_memory_valid_num"] = update_memory_valid_num + eval_results["overall_score"]["memory_update"]["update_memory_num"] = update_memory_num + + # Question-Answering Evaluation + correct_qa_num = 0 + hallucination_qa_num = 0 + omission_qa_num = 0 + qa_num = 0 + qa_valid_num = 0 + + for item in eval_results["question_answering_records"]: + item["is_valid"] = True + qa_num += 1 + + if item["result_type"] not in ["Correct", "Hallucination", "Omission"]: + item["is_valid"] = False + continue + + if item["result_type"] == "Correct": + correct_qa_num += 1 + elif item["result_type"] == "Hallucination": + hallucination_qa_num += 1 + elif item["result_type"] == "Omission": + omission_qa_num += 1 + + qa_valid_num += 1 + + if qa_num > 0: + eval_results["overall_score"]["question_answering"]["correct_qa_ratio(all)"] = correct_qa_num / qa_num + eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(all)"] = ( + hallucination_qa_num / qa_num + ) + eval_results["overall_score"]["question_answering"]["omission_qa_ratio(all)"] = omission_qa_num / qa_num + else: + eval_results["overall_score"]["question_answering"]["correct_qa_ratio(all)"] = 0 + eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(all)"] = 0 + eval_results["overall_score"]["question_answering"]["omission_qa_ratio(all)"] = 0 + + if qa_valid_num > 0: + eval_results["overall_score"]["question_answering"]["correct_qa_ratio(valid)"] = correct_qa_num / qa_valid_num + eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(valid)"] = ( + hallucination_qa_num / qa_valid_num + ) + eval_results["overall_score"]["question_answering"]["omission_qa_ratio(valid)"] = omission_qa_num / qa_valid_num + else: + eval_results["overall_score"]["question_answering"]["correct_qa_ratio(valid)"] = 0 + eval_results["overall_score"]["question_answering"]["hallucination_qa_ratio(valid)"] = 0 + eval_results["overall_score"]["question_answering"]["omission_qa_ratio(valid)"] = 0 + + eval_results["overall_score"]["question_answering"]["qa_valid_num"] = qa_valid_num + eval_results["overall_score"]["question_answering"]["qa_num"] = qa_num + + # Memory Type Accuracy + for item in eval_results["memory_integrity_records"]: + if "memory_integrity_score" not in item or "importance" not in item: + continue + score = 1 if item["memory_integrity_score"] == 2 else 0 + eval_results["overall_score"]["memory_type_accuracy"][item["memory_type"]]["memory_integrity_acc"] += score + eval_results["overall_score"]["memory_type_accuracy"][item["memory_type"]]["total_num"] += 1 + + for item in eval_results["memory_update_records"]: + if "memory_update_type" not in item or "importance" not in item: + continue + score = 1 if item["memory_update_type"] == "Correct" else 0 + eval_results["overall_score"]["memory_type_accuracy"][item["memory_type"]]["memory_update_acc"] += score + eval_results["overall_score"]["memory_type_accuracy"][item["memory_type"]]["total_num"] += 1 + + for key in eval_results["overall_score"]["memory_type_accuracy"]: + if eval_results["overall_score"]["memory_type_accuracy"][key]["total_num"] > 0: + total = eval_results["overall_score"]["memory_type_accuracy"][key]["total_num"] + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_integrity_acc"] = ( + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_integrity_acc"] / total + ) + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_update_acc"] = ( + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_update_acc"] / total + ) + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_acc"] = ( + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_integrity_acc"] + + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_update_acc"] + ) + else: + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_integrity_acc"] = 0 + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_update_acc"] = 0 + eval_results["overall_score"]["memory_type_accuracy"][key]["memory_acc"] = 0 + + return eval_results + + +# ==================== Main Pipeline ==================== + + +async def main_async( + data_path: str, + top_k: int = 20, + user_num: int = 1, + max_concurrency: int = 2, +): + """Main evaluation pipeline.""" + frame = "reme" + save_path = f"bench_results/{frame}/" + os.makedirs(save_path, exist_ok=True) + + output_file_stage1 = os.path.join(save_path, f"{frame}_eval_results.jsonl") + output_file_stage2 = os.path.join(save_path, f"{frame}_eval_stat_result.json") + + start_time = time.time() + await reme.vector_store.delete_all() + + # ==================== Stage 1: Data Processing ==================== + print("\n" + "=" * 80) + print("STAGE 1: PROCESSING DATA WITH ReMe") + print(f"Max Concurrency: {max_concurrency}") + print("=" * 80) + + tmp_dir = os.path.join(save_path, "tmp") + os.makedirs(tmp_dir, exist_ok=True) + + # Load all user data + user_data_list = list(iter_jsonl(data_path)) + total_users = min(len(user_data_list), user_num) + user_data_list = user_data_list[:total_users] + + print(f"Processing {total_users} users with max concurrency {max_concurrency}...") + + # Create semaphore to limit concurrency for Stage 1 + semaphore_stage1 = asyncio.Semaphore(max_concurrency) + + async def process_single_user_stage1(idx: int, user_data: dict): + """Process a single user in Stage 1 with semaphore control.""" + async with semaphore_stage1: + uuid = user_data['uuid'] + tmp_file = os.path.join(tmp_dir, f"{uuid}.json") + + if os.path.exists(tmp_file): + print(f"⚡ Skipping user {uuid} ({idx}/{total_users}) — cached result found.") + return {"uuid": uuid, "status": "cached", "path": tmp_file} + + print(f"[{idx}/{total_users}] Processing user {uuid}...") + result = await process_user_stage1(user_data, top_k, save_path) + print(f"[{idx}/{total_users}] ✅ Finished {uuid} ({result['status']})") + return result + + # Process users in parallel with controlled concurrency + tasks = [process_single_user_stage1(idx, user_data) for idx, user_data in enumerate(user_data_list, 1)] + await asyncio.gather(*tasks) + + # Combine all results into final output + with open(output_file_stage1, "w", encoding="utf-8") as f_out: + for file in os.listdir(tmp_dir): + if file.endswith(".json"): + file_path = os.path.join(tmp_dir, file) + with open(file_path, "r", encoding="utf-8") as f_in: + data = json.load(f_in) + f_out.write(json.dumps(data, ensure_ascii=False) + "\n") + + elapsed_stage1 = time.time() - start_time + print(f"\n✅ Stage 1 completed in {elapsed_stage1:.2f}s") + print(f"✅ Results saved to: {output_file_stage1}") + + # ==================== Stage 2: Evaluation ==================== + print("\n" + "=" * 80) + print("STAGE 2: EVALUATING MEMORY PERFORMANCE (Sequential)") + print("=" * 80) + + tmp_dir2 = os.path.join(save_path, "tmp2") + os.makedirs(tmp_dir2, exist_ok=True) + + start_stage2 = time.time() + + # Load all users and process sequentially + user_data_list = list(enumerate(iter_jsonl(output_file_stage1), 1)) + + for idx, user_data in user_data_list: + uuid = user_data["uuid"] + tmp_file = os.path.join(tmp_dir2, f"{uuid}.json") + + if os.path.exists(tmp_file): + print(f"⚡ Skipping user {uuid} ({idx}/{len(user_data_list)}) — cached result found.") + continue + + print(f"[{idx}/{len(user_data_list)}] Processing user {uuid}...") + t_user_result = await process_user_stage2(idx, user_data) + + with open(tmp_file, "w", encoding="utf-8") as f: + json.dump(t_user_result, f, ensure_ascii=False, indent=4) + + elapsed = time.time() - start_stage2 + print(f"[{idx}/{len(user_data_list)}] ✅ Finished user {uuid}, elapsed {elapsed:.2f}s.") + + # Calculate time consuming + add_dialogue_duration_time = 0 + search_memory_duration_time = 0 + + for user_data in iter_jsonl(output_file_stage1): + sessions = user_data["sessions"] + + for session in sessions: + if "add_dialogue_duration_ms" in session: + add_dialogue_duration_time += session["add_dialogue_duration_ms"] + + if "questions" in session: + for question in session["questions"]: + if "search_duration_ms" in question: + search_memory_duration_time += question["search_duration_ms"] + + add_dialogue_duration_time = add_dialogue_duration_time / 1000 / 60 + search_memory_duration_time = search_memory_duration_time / 1000 / 60 + + print("\n🔄 Aggregating all user results...") + + eval_results = { + "overall_score": { + "memory_integrity": {}, + "memory_accuracy": {}, + "memory_extraction_f1": 0, + "memory_update": {}, + "question_answering": {}, + "memory_type_accuracy": { + "Event Memory": { + "memory_integrity_acc": 0, + "memory_update_acc": 0, + "total_num": 0, + }, + "Persona Memory": { + "memory_integrity_acc": 0, + "memory_update_acc": 0, + "total_num": 0, + }, + "Relationship Memory": { + "memory_integrity_acc": 0, + "memory_update_acc": 0, + "total_num": 0, + }, + }, + "time_consuming": { + "add_dialogue_duration_time": add_dialogue_duration_time, + "search_memory_duration_time": search_memory_duration_time, + "total_duration_time": add_dialogue_duration_time + search_memory_duration_time, + }, + }, + "memory_integrity_records": [], + "memory_accuracy_records": [], + "memory_update_records": [], + "question_answering_records": [], + } + + for file_name in os.listdir(tmp_dir2): + if not file_name.endswith(".json"): + continue + user_file = os.path.join(tmp_dir2, file_name) + with open(user_file, "r", encoding="utf-8") as f: + user_result = json.load(f) + + eval_results["memory_accuracy_records"].extend(user_result.get("memory_accuracy_records", [])) + eval_results["memory_integrity_records"].extend(user_result.get("memory_integrity_records", [])) + eval_results["memory_update_records"].extend(user_result.get("memory_update_records", [])) + eval_results["question_answering_records"].extend(user_result.get("question_answering_records", [])) + + eval_results = aggregate_eval_results(eval_results) + + with open(output_file_stage2, "w", encoding="utf-8") as f: + json.dump(eval_results, f, ensure_ascii=False, indent=4) + + elapsed_total = time.time() - start_time + print(f"\n✅ All done in {elapsed_total:.2f}s. Results saved to {output_file_stage2}") + + # Print summary + print("\n" + "=" * 80) + print("EVALUATION SUMMARY") + print("=" * 80) + print("\n📊 Memory Integrity:") + print(f" - Recall (all): {eval_results['overall_score']['memory_integrity'].get('recall(all)', 0):.4f}") + print(f" - Recall (valid): {eval_results['overall_score']['memory_integrity'].get('recall(valid)', 0):.4f}") + print(f" - Weighted Recall (all): " + f"{eval_results['overall_score']['memory_integrity'].get('weighted_recall(all)', 0):.4f}") + + print(f"\n📊 Memory Accuracy:") + print(f" - Target Accuracy (all): {eval_results['overall_score']['memory_accuracy'].get('target_accuracy(all)', 0):.4f}") + print(f" - Target Accuracy (valid): {eval_results['overall_score']['memory_accuracy'].get('target_accuracy(valid)', 0):.4f}", + ) + print( + f" - Weighted Accuracy (all): {eval_results['overall_score']['memory_accuracy'].get('weighted_accuracy(all)', 0):.4f}", + ) + + print(f"\n📊 Memory Extraction F1: {eval_results['overall_score']['memory_extraction_f1']:.4f}") + + print(f"\n📊 Memory Update:") + print( + f" - Correct (all): {eval_results['overall_score']['memory_update'].get('correct_update_memory_ratio(all)', 0):.4f}", + ) + print( + f" - Hallucination (all): {eval_results['overall_score']['memory_update'].get('hallucination_update_memory_ratio(all)', 0):.4f}", + ) + print( + f" - Omission (all): {eval_results['overall_score']['memory_update'].get('omission_update_memory_ratio(all)', 0):.4f}", + ) + + print(f"\n📊 Question Answering:") + print( + f" - Correct (all): {eval_results['overall_score']['question_answering'].get('correct_qa_ratio(all)', 0):.4f}", + ) + print( + f" - Hallucination (all): {eval_results['overall_score']['question_answering'].get('hallucination_qa_ratio(all)', 0):.4f}", + ) + print( + f" - Omission (all): {eval_results['overall_score']['question_answering'].get('omission_qa_ratio(all)', 0):.4f}", + ) + + print(f"\n⏱️ Time Consuming:") + print(f" - Add Dialogue: {add_dialogue_duration_time:.2f} min") + print(f" - Search Memory: {search_memory_duration_time:.2f} min") + print(f" - Total: {add_dialogue_duration_time + search_memory_duration_time:.2f} min") + print("=" * 80) + + +def main( + data_path: str, + top_k: int = 20, + user_num: int = 1, + max_concurrency: int = 2, +): + """Synchronous entry point.""" + asyncio.run(main_async(data_path, top_k, user_num, max_concurrency)) + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser(description="Complete evaluation for ReMe on HaluMem benchmark") + parser.add_argument( + "--data_path", + type=str, + required=True, + help="Path to HaluMem data file (e.g., HaluMem-medium.jsonl)", + ) + parser.add_argument( + "--top_k", + type=int, + default=20, + help="Number of top memories to retrieve (default: 20)", + ) + parser.add_argument( + "--user_num", + type=int, + default=1, + help="Number of users to evaluate (default: 1)", + ) + parser.add_argument( + "--max_concurrency", + type=int, + default=2, + help="Maximum concurrency for stage 1 processing (default: 2)", + ) + args = parser.parse_args() + + main( + data_path=args.data_path, + top_k=args.top_k, + user_num=args.user_num, + max_concurrency=args.max_concurrency, + ) diff --git a/bench/halumem/eval_reme_simple.py b/bench/halumem/eval_reme_simple.py new file mode 100644 index 00000000..a7ac6964 --- /dev/null +++ b/bench/halumem/eval_reme_simple.py @@ -0,0 +1,516 @@ +""" +Simplified evaluation script for ReMe on HaluMem benchmark - Question Answering only. + +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 + +Usage: + python bench/halumem/eval_reme_simple.py --data_path /Users/yuli/workspace/HaluMem/data/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 datetime import datetime, timezone + +from loguru import logger + +from eval_tools import ( + _PROMPTS, + evaluation_for_question, + evaluation_for_question2, +) +from llms import llm_request +from reme_ai.core.enumeration import MemoryType +from reme_ai.core.schema import MemoryNode +from reme_ai.reme import ReMe + +# Initialize ReMe +reme: ReMe = ReMe() + + +def extract_user_name(persona_info: str): + """Extract user name from persona info.""" + match = re.search(r"Name:\s*(.*?); Gender:", persona_info) + if match: + username = match.group(1).strip() + return username + else: + raise ValueError("No name found.") + + +def iter_jsonl(file_path: str): + """Iterate over lines in a JSONL file.""" + with open(file_path, "r", encoding="utf-8") as f: + for line in f: + line = line.strip() + if line: + yield json.loads(line) + + +# ==================== 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) + + # Format the context + context = f"User: {user_id}\nMemories:\n{memories}" + + # 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) + + 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 = [ + { + "role": turn["role"], + "content": turn["content"], + "time_created": datetime.strptime(turn["timestamp"], "%b %d, %Y, %H:%M:%S") + .replace(tzinfo=timezone.utc) + .strftime("%Y-%m-%d %H:%M:%S"), + } + for turn in dialogue + ] + + # Process in batches + result = [] + 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, + ) + 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 + }) + total_duration_ms += duration_ms + + duration_ms = total_duration_ms + + # Extract memory content + memories = [] + for memory_node in result: + if isinstance(memory_node, MemoryNode) and memory_node.memory_type is not MemoryType.HISTORY: + memories.append(memory_node.content) + + if session.get("is_generated_qa_session", False): + new_session["add_dialogue_duration_ms"] = duration_ms + new_session["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 + + # 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 + + new_session["questions"] = [] + + for qa in session["questions"]: + response, agent_messages, success, duration_ms = await search_memory_async( + query=qa["question"], + user_id=user_name, + top_k=top_k_value, + ) + + new_qa = copy.deepcopy(qa) + new_qa["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) + + 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} + + +# ==================== Evaluation Aggregation ==================== + + +def aggregate_eval_results(eval_results): + """Aggregate evaluation results and compute metrics (QA only).""" + + # 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 + + +# ==================== 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}...") + + # 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") + + if os.path.exists(tmp_file): + print(f"⚡ Skipping user {uuid} ({idx}/{total_users}) — cached result found.") + return {"uuid": uuid, "status": "cached", "path": tmp_file} + + print(f"[{idx}/{total_users}] Processing user {uuid}...") + result = await process_user_stage1(user_data, top_k, save_path) + print(f"[{idx}/{total_users}] ✅ Finished {uuid} ({result['status']})") + return result + + # Process users in parallel with controlled concurrency + tasks = [process_single_user(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, + }, + }, + "question_answering_records": [], + } + + # 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) + + +def main( + data_path: str, + top_k: int = 20, + user_num: int = 1, + max_concurrency: int = 2, +): + """Synchronous entry point.""" + asyncio.run(main_async(data_path, top_k, user_num, max_concurrency)) + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser(description="Simplified evaluation for ReMe on HaluMem benchmark (QA only)") + parser.add_argument( + "--data_path", + type=str, + required=True, + help="Path to HaluMem data file (e.g., HaluMem-medium.jsonl)", + ) + parser.add_argument( + "--top_k", + type=int, + default=20, + help="Number of top memories to retrieve (default: 20)", + ) + parser.add_argument( + "--user_num", + type=int, + default=1, + help="Number of users to evaluate (default: 1)", + ) + parser.add_argument( + "--max_concurrency", + type=int, + default=2, + help="Maximum concurrency for 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, + ) diff --git a/bench/halumem/eval_tools.py b/bench/halumem/eval_tools.py new file mode 100644 index 00000000..bf68ce2d --- /dev/null +++ b/bench/halumem/eval_tools.py @@ -0,0 +1,133 @@ +"""Evaluation tools for ReMe HaluMem benchmark.""" + +from pathlib import Path + +import yaml + +from llms import llm_request_for_json + +# Load prompts from YAML file +_YAML_PATH = Path(__file__).parent / "halumem.yaml" +with open(_YAML_PATH, "r", encoding="utf-8") as f: + _PROMPTS = yaml.safe_load(f) + + +async def evaluation_for_memory_integrity( + extract_memories: str, + target_memory: str, +): + """ + Memory Integrity Evaluation + extract_memories: A formatted string concatenating all memory points extracted by the memory system under evaluation. + target_memory: The target key memory point. + """ + + prompt = _PROMPTS["EVALUATION_PROMPT_FOR_MEMORY_INTEGRITY"].format( + memories=extract_memories, + expected_memory_point=target_memory, + ) + + result = await llm_request_for_json(prompt) + + return result + + +async def evaluation_for_memory_accuracy( + dialogue: str, + golden_memories: str, + candidate_memory: str, +): + """ + Memory Accuracy Evaluation + dialogue: The complete human-machine dialogue record. + golden_memories: The core memory points for this dialogue segment in the evaluation set (the correct reference memories). + candidate_memory: A specific memory point extracted by the memory system being evaluated. + """ + + prompt = _PROMPTS["EVALUATION_PROMPT_FOR_MEMORY_ACCURACY"].format( + dialogue=dialogue, + golden_memories=golden_memories, + candidate_memory=candidate_memory, + ) + + result = await llm_request_for_json(prompt) + + return result + + +async def evaluation_for_update_memory( + extract_memories: str, + target_update_memory: str, + original_memory: str, +): + """ + Memory Update Evaluation + extract_memories: A formatted string concatenating all memory points extracted by the memory system under evaluation. + target_update_memory: The target updated memory point. + original_memory: str: A formatted string concatenating all original memory points corresponding to the target updated memory point (i.e., all memories before the update). + """ + + prompt = _PROMPTS["EVALUATION_PROMPT_FOR_UPDATE_MEMORY"].format( + memories=extract_memories, + updated_memory=target_update_memory, + original_memory=original_memory, + ) + + result = await llm_request_for_json(prompt) + + return result + + +async def evaluation_for_question( + question: str, + reference_answer: str, + key_memory_points: str, + response: str, +): + """ + Question-Answering Evaluation + question: The question string to be evaluated. + reference_answer: The reference (gold-standard) answer. + key_memory_points: The memory points used to derive the reference answer. + response: The answer produced by the memory system. + """ + + prompt = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION"].format( + question=question, + reference_answer=reference_answer, + key_memory_points=key_memory_points, + response=response, + ) + + result = await llm_request_for_json(prompt) + + return result + + +async def evaluation_for_question2( + question: str, + reference_answer: str, + key_memory_points: str, + response: str, + dialogue: str, +): + """ + Question-Answering Evaluation with Dialogue Context (Version 2) + question: The question string to be evaluated. + reference_answer: The reference (gold-standard) answer. + key_memory_points: The memory points used to derive the reference answer. + response: The answer produced by the memory system. + dialogue: The formatted dialogue history (role, content, time_created). + """ + + prompt = _PROMPTS["EVALUATION_PROMPT_FOR_QUESTION2"].format( + question=question, + reference_answer=reference_answer, + key_memory_points=key_memory_points, + response=response, + dialogue=dialogue, + ) + + result = await llm_request_for_json(prompt) + + return result diff --git a/bench/halumem/halumem.yaml b/bench/halumem/halumem.yaml new file mode 100644 index 00000000..60f7cbda --- /dev/null +++ b/bench/halumem/halumem.yaml @@ -0,0 +1,494 @@ +TEMPLATE_MEMOS: | + Memories for user {user_id}: + {memories} + +PROMPT_MEMZERO: | + You are an intelligent memory assistant tasked with retrieving accurate information from conversation memories. + + # CONTEXT: + You have access to memories from two speakers in a conversation. These memories contain + timestamped information that may be relevant to answering the question. + + # INSTRUCTIONS: + 1. Carefully analyze all provided memories from both speakers + 2. Pay special attention to the timestamps to determine the answer + 3. If the question asks about a specific event or fact, look for direct evidence in the memories + 4. If the memories contain contradictory information, prioritize the most recent memory + 5. If there is a question about time references (like "last year", "two months ago", etc.), + calculate the actual date based on the memory timestamp. For example, if a memory from + 4 May 2022 mentions "went to India last year," then the trip occurred in 2021. + 6. Always convert relative time references to specific dates, months, or years. For example, + convert "last year" to "2022" or "two months ago" to "March 2023" based on the memory + timestamp. Ignore the reference while answering the question. + 7. Focus only on the content of the memories from both speakers. Do not confuse character + names mentioned in memories with the actual users who created those memories. + 8. The answer should be less than 5-6 words. + + # APPROACH (Think step by step): + 1. First, examine all memories that contain information related to the question + 2. Examine the timestamps and content of these memories carefully + 3. Look for explicit mentions of dates, times, locations, or events that answer the question + 4. If the answer requires calculation (e.g., converting relative time references), show your work + 5. Formulate a precise, concise answer based solely on the evidence in the memories + 6. Double-check that your answer directly addresses the question asked + 7. Ensure your final answer is specific and avoids vague time references + + {context} + + Question: {question} + + Answer: + +PROMPT_ZEP: | + You are an intelligent memory assistant tasked with retrieving accurate information from conversation memories. + + # CONTEXT: + You have access to memories from a conversation. These memories contain + timestamped information that may be relevant to answering the question. + + # INSTRUCTIONS: + 1. Carefully analyze all provided memories + 2. Pay special attention to the timestamps to determine the answer + 3. If the question asks about a specific event or fact, look for direct evidence in the memories + 4. If the memories contain contradictory information, prioritize the most recent memory + 5. If there is a question about time references (like "last year", "two months ago", etc.), + calculate the actual date based on the memory timestamp. For example, if a memory from + 4 May 2022 mentions "went to India last year," then the trip occurred in 2021. + 6. Always convert relative time references to specific dates, months, or years. For example, + convert "last year" to "2022" or "two months ago" to "March 2023" based on the memory + timestamp. Ignore the reference while answering the question. + 7. Focus only on the content of the memories. Do not confuse character + names mentioned in memories with the actual users who created those memories. + 8. The answer should be less than 5-6 words. + + # APPROACH (Think step by step): + 1. First, examine all memories that contain information related to the question + 2. Examine the timestamps and content of these memories carefully + 3. Look for explicit mentions of dates, times, locations, or events that answer the question + 4. If the answer requires calculation (e.g., converting relative time references), show your work + 5. Formulate a precise, concise answer based solely on the evidence in the memories + 6. Double-check that your answer directly addresses the question asked + 7. Ensure your final answer is specific and avoids vague time references + + Context: + + {context} + + Question: {question} + Answer: + +PROMPT_MEMOS: | + You are a knowledgeable and helpful AI assistant. + + # CONTEXT: + You have access to memories from two speakers in a conversation. These memories contain + timestamped information that may be relevant to answering the question. + + # INSTRUCTIONS: + 1. Carefully analyze all provided memories. Synthesize information across different entries if needed to form a complete answer. + 2. Pay close attention to the timestamps to determine the answer. If memories contain contradictory information, the **most recent memory** is the source of truth. + 3. If the question asks about a specific event or fact, look for direct evidence in the memories. + 4. Your answer must be grounded in the memories. However, you may use general world knowledge to interpret or complete information found within a memory (e.g., identifying a landmark mentioned by description). + 5. If the question involves time references (like "last year", "two months ago", etc.), you **must** calculate the actual date based on the memory's timestamp. For example, if a memory from 4 May 2022 mentions "went to India last year," then the trip occurred in 2021. + 6. Always convert relative time references to specific dates, months, or years in your final answer. + 7. Do not confuse character names mentioned in memories with the actual users who created them. + 8. The answer must be brief (under 5-6 words) and direct, with no extra description. + + # APPROACH (Think step by step): + 1. First, examine all memories that contain information related to the question. + 2. Synthesize findings from multiple memories if a single entry is insufficient. + 3. Examine timestamps and content carefully, looking for explicit dates, times, locations, or events. + 4. If the answer requires calculation (e.g., converting relative time references), perform the calculation. + 5. Formulate a precise, concise answer based on the evidence from the memories (and allowed world knowledge). + 6. Double-check that your answer directly addresses the question asked and adheres to all instructions. + 7. Ensure your final answer is specific and avoids vague time references. + + {context} + + Question: {question} + + Answer: + +PROMPT_MEMOBASE: | + You are an intelligent memory assistant tasked with retrieving accurate information from conversation memories. + + # CONTEXT: + You have access to memories from two speakers in a conversation. These memories contain + timestamped information that may be relevant to answering the question. + + # INSTRUCTIONS: + 1. Carefully analyze all provided memories from both speakers + 2. Pay special attention to the timestamps to determine the answer + 3. If the question asks about a specific event or fact, look for direct evidence in the memories + 4. If the memories contain contradictory information, prioritize the most recent memory + 5. If there is a question about time references (like "last year", "two months ago", etc.), calculate the actual date based on the memory timestamp. For example, if a memory from 4 May 2022 mentions "went to India last year," then the trip occurred in 2021. + 6. Always convert relative time references to specific dates, months, or years. For example, convert "last year" to "2022" or "two months ago" to "March 2023" based on the memory timestamp. Ignore the reference while answering the question. + 7. Focus only on the content of the memories from both speakers. Do not confuse character names mentioned in memories with the actual users who created those memories. + 8. The answer should be less than 5-6 words. + + # APPROACH (Think step by step): + 1. First, examine all memories that contain information related to the question + 2. Examine the timestamps and content of these memories carefully + 3. Look for explicit mentions of dates, times, locations, or events that answer the question + 4. If the answer requires calculation (e.g., converting relative time references), show your work + 5. Formulate a precise, concise answer based solely on the evidence in the memories + 6. Double-check that your answer directly addresses the question asked + 7. Ensure your final answer is specific and avoids vague time references + + {context} + + Question: {question} + + Answer: + + +EVALUATION_PROMPT_FOR_MEMORY_INTEGRITY: | + You are a strict **"Memory Integrity" evaluator**. + Your core task is to assess whether an AI memory system has **missed any key memory points** after processing a conversation. This evaluation measures the system’s **memory integrity**, i.e., its ability to resist **amnesia** or **omission**. + + # Evaluation Context & Data: + + 1. **Extracted Memories:** + These are all the memory items actually extracted by the memory system. + {memories} + + 2. **Expected Memory Point:** + The key memory point that *should* have been extracted. + {expected_memory_point} + + # Evaluation Instructions: + + 1. For each **Expected Memory Point**, search within the **Extracted Memories** list for corresponding or related information. Ignore unrelated items. + 2. Based on the following scoring rubric, rate how well the memory system captured the **Expected Memory Point** and provide a detailed explanation. + + # Scoring Rubric: + + * **2:** Fully covered or implied. + One or more items in “Extracted Memories” fully cover or logically imply all information in the “Expected Memory Point.” + + * **1:** Partially covered or mentioned. + Some information in “Extracted Memories” mentions part of the “Expected Memory Point,” but key information is missing, inaccurate, or slightly incorrect. + + * **0:** Not mentioned or incorrect. + “Extracted Memories” contains no mention of the “Expected Memory Point,” or the corresponding information is entirely wrong. + + # Scoring Notes: + + * For **compound Expected Memory Points** (with multiple elements such as person/event/time/location/preference, etc.): + + * All elements correct → **2 points** + * Some elements correct / uncertain → **1 point** + * Key elements missing or wrong → **0 points** + + * Semantic matching is acceptable; exact wording is **not** required. + + * If “Extracted Memories” contains **conflicting information**, assign the **best possible coverage score** and mention the conflict in your reasoning. + + * Extra or stylistically different memories do **not** reduce the score; only the coverage of the **Expected Memory Point** matters. + + * For uncertain wording (“might,” “probably,” “tends to,” etc.): + + * If the Expected Memory Point is a definite statement, usually assign **1 point**. + + * If critical fields (e.g., time, entity name, relationship) are partly wrong but others match → **1 point**. + + * If all key fields are wrong or missing → **0 points**. + + # Output Format: + + Please output your result in the following JSON format: + + ```json + {{ + "reasoning": "Provide a concise justification for the score", + "score": "2|1|0" + }} + ``` + +EVALUATION_PROMPT_FOR_MEMORY_ACCURACY: | + You are a **Dialogue Memory Accuracy Evaluator.** Your task is to evaluate the **accuracy** of a memory extracted by an AI memory system, based on three given inputs: the dialogue content, the *target (gold)* memory points (the correct annotated memories), and the *candidate* memory to be evaluated. The goal is to output a **structured evaluation result**. + + # Input Content + + * **Dialogue:** + {dialogue} + + * **Golden Memories (Target Memory Points):** + The correct memory points pre-annotated for this dialogue in the evaluation dataset. + {golden_memories} + + * **Candidate Memory:** + The memory extracted by the system to be evaluated. + {candidate_memory} + + # Evaluation Principles and Definitions + + ### 1) Support / Entailment + + * An **information point** (atomic fact) in the candidate memory is considered *supported* if it can be directly stated or semantically entailed (via synonym, paraphrase, or equivalent expression) by the *Dialogue* or *Golden Memories*. + * Only the given dialogue and golden memories can be used for judgment — **no external knowledge** or assumptions are allowed. + Any information not appearing in or inferable from these two sources is considered *unsupported*. + * Pay careful attention to **negation**, **quantities**, **time**, and **subjects**. + If the candidate statement contradicts the dialogue or golden memories, it is considered a **conflict**. + + ### 2) Memory Accuracy Score (integer: 0 / 1 / 2) + + * **2 points:** Every information point in the candidate memory is supported by the dialogue or golden memories, with **no contradictions or hallucinations**. + * **1 point:** The candidate memory is *partially correct* (at least one supported information point) but also includes *unsupported* or *contradictory* content. + * **0 points:** The candidate memory is **entirely unsupported or contradictory** to the sources (i.e., a “hallucinated memory”). + + > Note: + > + > * If a candidate memory contains multiple information points, **any unsupported or contradictory element** prevents a full score (2). + > * If both supported and unsupported/conflicting content appear, assign a score of **1**. + + ### 3) Inclusion in Golden Memories (Boolean field-level judgment) + + **Definition:** + + * **Atomic information point:** the smallest factual unit in the candidate memory (e.g., *name = Li Si*, *age = 25*, *location = Beijing*, *preference = coffee*, *budget ≤ 2000*, *meeting_time = Wednesday 10:00*, *tool = Zoom*, etc.). + * **Field / Slot:** the semantic dimension of an information point (e.g., *name*, *age*, *residence*, *food preference*, *budget*, *meeting time*, *meeting tool*, etc.). + + **Judgment Rules (independent of correctness):** + + * **true:** + Every atomic information point in the candidate memory has a corresponding **field** in the golden memories (allowing for synonyms, paraphrases, or equivalent expressions; ignore value, polarity, or quantity differences). + + * Note: A single field in the gold list may match multiple candidate points (e.g., multiple “drink preference” facts can be covered by one “drink preference” field in gold). + * **false:** + If **any** atomic information point’s field in the candidate memory cannot be found in the golden memories, mark as *false*. + + **Important Notes:** + + * Field matching is restricted to fields that are **explicitly present or semantically recognizable** in the golden memories — no external knowledge may be used to expand the field set. + * Differences in **values** (e.g., “Zhang San” vs. “Li Si”), **polarity** (like/dislike), or **exact number/time** do **not** affect this Boolean judgment. + + # Evaluation Procedure + + For each candidate memory: + + 1. **Decompose** it into atomic information points (e.g., name, number, location, preference). + 2. For each information point, **search** the dialogue and golden memories for supporting or contradictory evidence. + 3. Assign the **accuracy_score** (0 / 1 / 2) according to the rules above. + 4. Determine **is_included_in_golden_memories (true/false)**: + + * Identify each information point’s field; + * If *all* fields exist in the golden memories, mark as *true*; otherwise, *false*. + 5. Provide a **concise Chinese explanation** in `"reason"`, citing key evidence (short excerpts allowed), and clearly state any unsupported or contradictory parts if applicable. + + # Output Format (strictly required) + + Output **only one JSON object**, with the following three fields: + + * `"accuracy_score"`: `"0"` or `"1"` or `"2"` + * `"is_included_in_golden_memories"`: `"true"` or `"false"` + * `"reason"`: `"brief explanation in Chinese"` + + Do **not** include any other text, explanation, or fields. + Do **not** include the candidate memory text inside the JSON. + + Please output **only** the following JSON (in a code block): + + ```json + {{ + "accuracy_score": "2 | 1 | 0", + "is_included_in_golden_memories": "true | false", + "reason": "Brief explanation in Chinese" + }} + ``` + +EVALUATION_PROMPT_FOR_UPDATE_MEMORY: | + Your task is to **evaluate the update accuracy** of an AI memory system. + Based on the information provided below, determine whether the system-generated **“Generated Memories”** correctly **includes** the **Target Memory for Update**. + + # Background Information + + The following information is provided for evaluation: + + 1. **Generated Memories:** + This is the list of memory points generated by the system after the current dialogue. + {memories} + + 2. **Target Memory for Update:** + This is the correct, updated version of the memory point that should have been produced — the one we focus on in this evaluation. + {updated_memory} + + 3. **Original Memory Content:** + This is the original version of the target memory before the update. + {original_memory} + + # Evaluation Criteria + + Please make your judgment **strictly based on the content update of the “Target Memory for Update.”** + Use the following categories: + + ### Correct Update + + * **Generated Memories** **contains all information points** from the “Target Memory for Update,” accurately and completely reflecting the intended update. + * **Key fields** (e.g., date, time, values, proper nouns, etc.) must match exactly. + * The **original memory** is effectively replaced or marked as outdated. + * Synonymous or slightly rephrased expressions are acceptable. + + ### Hallucinated Update + + * **Factual error:** The **Generated Memories** includes a new memory related to the “Target Memory for Update,” but its content contains factual mistakes or contradictions compared to the correct update. + + ### Omitted Update + + * **Completely omitted:** The **Generated Memories** contains no new memory related to the “Target Memory for Update.” + * **Partially omitted:** A related new memory was generated in **Generated Memories**, but it **misses key information** that should have been included. + + ### Other + + Used for update failures that do **not clearly fall** into the above categories of “Hallucination” or “Omission.” + + # Output Requirements + + Please return your evaluation strictly in the following JSON format and provide a concise explanation. + + ```json + {{ + "reason": "Briefly explain your reasoning here and why it fits this category.", + "evaluation_result": "Correct | Hallucination | Omission | Other" + }} + ``` + +EVALUATION_PROMPT_FOR_QUESTION: | + You are an **evaluation expert for AI memory system question answering**. + Based **only** on the provided **“Question”**, **“Reference Answer”**, and **“Key Memory Points”** (the essential facts needed to derive the reference answer), strictly evaluate the **accuracy** of the **“Memory System Response.”** Classify it as one of **“Correct”**, **“Hallucination”**, or **“Omission.”** Do **not** use any external knowledge or subjective inference. Finally, output your judgment **strictly** in the specified JSON format. + + # Evaluation Criteria + + ## Answer Type Classification + + ### 1. Correct + + * The “Memory System Response” accurately answers the “Question,” and its content is **semantically equivalent** to the “Reference Answer.” + * It contains **no contradictions** with the “Key Memory Points” or “Reference Answer.” + * It introduces **no unsupported details** beyond the “Key Memory Points” that could alter the conclusion. + * Synonyms, paraphrasing, and reasonable summarization are acceptable. + + ### 2. Hallucination + + * The “Memory System Response” includes information or facts that **contradict or are inconsistent** with the “Reference Answer” or the “Key Memory Points.” + * When the “Reference Answer” is labeled as *unknown/uncertain*, yet the response provides a specific verifiable fact or conclusion. + * Extra irrelevant information that does **not change** the conclusion is **not** considered hallucination by itself; however, if it **changes or misleads** the conclusion, or **contradicts** the “Key Memory Points,” it should be judged as a **Hallucination**. + + ### 3. Omission + + * The response is **incomplete** compared to the “Reference Answer.” + * It explicitly states “don’t know,” “can’t remember,” or “no related memory,” even though relevant information exists in the “Key Memory Points.” + * For multi-element questions, **all elements must be correct and present**; omission of **any** element is considered an **Omission**. + + ## Priority Rules (Conflict Handling) + + * If the response contains **both missing necessary information** and **fabricated/contradictory information**, classify it as **Hallucination**. + * If there is **no fabrication/contradiction** but some necessary information is missing, classify it as **Omission**. + * Only when the meaning is **fully equivalent** to the reference answer should it be classified as **Correct**. + + ## Detailed Guidelines and Tolerance + + * Equivalent expressions of numbers, times, and units are acceptable, but the **numerical values themselves must not differ**. + * For multi-element questions, **all elements must be complete and accurate**; missing any element counts as **Omission**. + * If the reference answer is *“unknown / cannot be determined”* and the system provides a definite fact, that is a **Hallucination**. + If the system also answers *“unknown”* (without guessing), it may be **Correct**. + * The evaluation must rely **only** on the *Reference Answer*, *Key Memory Points*, and *System Response* — no external context, world knowledge, or speculative reasoning is allowed. + + # Information for Evaluation + + * **Question:** + {question} + + * **Reference Answer:** + {reference_answer} + + * **Key Memory Points:** + {key_memory_points} + + * **Memory System Response:** + {response} + + # Output Requirements + + Please provide your evaluation result **strictly** in the JSON format below. + Do **not** add any extra explanation or comments outside the JSON block. + + ```json + {{ + "reasoning": "Provide a concise and traceable evaluation rationale: first compare the system’s response with the Key Memory Points (which were correctly used, which were missing, and whether there was any fabrication/contradiction), then assess its consistency with the Reference Answer, and finally state the classification basis.", + "evaluation_result": "Correct | Hallucination | Omission" + }} + ``` + + +EVALUATION_PROMPT_FOR_QUESTION2: | + You are an **evaluation expert for AI memory system question answering**. + + **Dialogue:** + {dialogue} + + Based **only** on the provided **“Question”**, **“Reference Answer”**, and **“Key Memory Points”** (the essential facts needed to derive the reference answer), strictly evaluate the **accuracy** of the **“Memory System Response.”** Classify it as one of **“Correct”**, **“Hallucination”**, or **“Omission.”** Do **not** use any external knowledge or subjective inference. Finally, output your judgment **strictly** in the specified JSON format. + + # Evaluation Criteria + + ## Answer Type Classification + + ### 1. Correct + + * The “Memory System Response” accurately answers the “Question,” and its content is **semantically equivalent** to the “Reference Answer.” + * It contains **no contradictions** with the “Key Memory Points” or “Reference Answer.” + * It introduces **no unsupported details** beyond the “Key Memory Points” that could alter the conclusion. + * Synonyms, paraphrasing, and reasonable summarization are acceptable. + + ### 2. Hallucination + + * The “Memory System Response” includes information or facts that **contradict or are inconsistent** with the “Reference Answer” or the “Key Memory Points.” + * When the “Reference Answer” is labeled as *unknown/uncertain*, yet the response provides a specific verifiable fact or conclusion. + * Extra irrelevant information that does **not change** the conclusion is **not** considered hallucination by itself; however, if it **changes or misleads** the conclusion, or **contradicts** the “Key Memory Points,” it should be judged as a **Hallucination**. + + ### 3. Omission + + * The response is **incomplete** compared to the “Reference Answer.” + * It explicitly states “don’t know,” “can’t remember,” or “no related memory,” even though relevant information exists in the “Key Memory Points.” + * For multi-element questions, **all elements must be correct and present**; omission of **any** element is considered an **Omission**. + + ## Priority Rules (Conflict Handling) + + * If the response contains **both missing necessary information** and **fabricated/contradictory information**, classify it as **Hallucination**. + * If there is **no fabrication/contradiction** but some necessary information is missing, classify it as **Omission**. + * Only when the meaning is **fully equivalent** to the reference answer should it be classified as **Correct**. + + ## Detailed Guidelines and Tolerance + + * Equivalent expressions of numbers, times, and units are acceptable, but the **numerical values themselves must not differ**. + * For multi-element questions, **all elements must be complete and accurate**; missing any element counts as **Omission**. + * If the reference answer is *“unknown / cannot be determined”* and the system provides a definite fact, that is a **Hallucination**. + If the system also answers *“unknown”* (without guessing), it may be **Correct**. + * The evaluation must rely **only** on the *Reference Answer*, *Key Memory Points*, and *System Response* — no external context, world knowledge, or speculative reasoning is allowed. + + # Information for Evaluation + + * **Question:** + {question} + + * **Reference Answer:** + {reference_answer} + + * **Key Memory Points:** + {key_memory_points} + + * **Memory System Response:** + {response} + + # Output Requirements + + Please provide your evaluation result **strictly** in the JSON format below. + Do **not** add any extra explanation or comments outside the JSON block. + + ```json + {{ + "reasoning": "Provide a concise and traceable evaluation rationale: first compare the system’s response with the Key Memory Points (which were correctly used, which were missing, and whether there was any fabrication/contradiction), then assess its consistency with the Reference Answer, and finally state the classification basis.", + "evaluation_result": "Correct | Hallucination | Omission" + }} + ``` + """ \ No newline at end of file diff --git a/bench/halumem/llms.py b/bench/halumem/llms.py new file mode 100644 index 00000000..a6a06218 --- /dev/null +++ b/bench/halumem/llms.py @@ -0,0 +1,87 @@ +import asyncio +import json +import logging +import re + +from tenacity import retry, stop_after_attempt, wait_random_exponential, before_sleep_log + +from reme_ai.core.schema import Message +from reme_ai.core.utils import load_env +from reme_ai.reme import ReMe + +logger = logging.getLogger(__name__) + +load_env() + +WAIT_TIME_LOWER = 1 +WAIT_TIME_UPPER = 60 +RETRY_TIMES = 5 + +# Use ReMe singleton's LLM instead of creating a separate instance +reme = ReMe() + +@retry( + wait=wait_random_exponential(min=WAIT_TIME_LOWER, max=WAIT_TIME_UPPER), + stop=stop_after_attempt(3), + reraise=True, + before_sleep=before_sleep_log(logger, logging.WARNING), +) +async def llm_request(prompt, model_name: str = "qwen3-max", **kwargs) -> str: + """Make an LLM request using ReMe's LLM with optional model override. + + Args: + prompt: The prompt to send to the LLM + model_name: Optional model name to override the default model (default: "qwen3-max") + **kwargs: Additional arguments to pass to the chat method + + Returns: + The assistant's response content + """ + assistant_message = await reme.llm.chat( + messages=[ + Message( + **{ + "role": "user", + "content": prompt, + }, + ), + ], + model_name=model_name, + **kwargs, + ) + return assistant_message.content + + +@retry( + wait=wait_random_exponential(min=WAIT_TIME_LOWER, max=WAIT_TIME_UPPER), + stop=stop_after_attempt(RETRY_TIMES), + reraise=True, + before_sleep=before_sleep_log(logger, logging.WARNING), +) +async def llm_request_for_json(prompt, model_name: str = "qwen3-max", **kwargs): + """Make an LLM request expecting JSON response using ReMe's LLM. + + Args: + prompt: The prompt to send to the LLM + model_name: Optional model name to override the default model (default: "qwen3-max") + **kwargs: Additional arguments to pass to the chat method + + Returns: + Parsed JSON object from the LLM response + + Raises: + ValueError: If no JSON block is found in the model output + """ + content = await llm_request(prompt, model_name=model_name, **kwargs) + + match = re.search(r"```json\s*(\{.*?\})\s*```", content, re.DOTALL) + if not match: + raise ValueError(f"No JSON block found in model output: {content}") + + json_str = match.group(1).strip() + return json.loads(json_str) + + +if __name__ == "__main__": + r = asyncio.run(llm_request_for_json('hello? answer in ```json\n{"answer": "..."}```')) + print(r) diff --git a/reme_ai/core/application.py b/reme_ai/core/application.py index c0561705..29eb846e 100644 --- a/reme_ai/core/application.py +++ b/reme_ai/core/application.py @@ -69,8 +69,6 @@ class Application: self._update_env("REME_EMBEDDING_API_KEY", embedding_api_key) self._update_env("REME_EMBEDDING_BASE_URL", embedding_api_base) - init_logger() - # Use default parser if not provided parser_class = parser if parser is not None else PydanticConfigParser self.parser = parser_class(ServiceConfig) @@ -87,6 +85,9 @@ class Application: C.service_config = service_config + if C.service_config.init_logger: + init_logger() + if llm: C.update_section_config("llm", **llm) if embedding_model: diff --git a/reme_ai/core/config/default.yaml b/reme_ai/core/config/default.yaml index 0b7276f7..16b03f60 100644 --- a/reme_ai/core/config/default.yaml +++ b/reme_ai/core/config/default.yaml @@ -16,11 +16,15 @@ llm: default: backend: openai model_name: qwen3-30b-a3b-instruct-2507 + max_rps: 6 + rps_window: 10 qwen3_max_instruct: backend: openai model_name: qwen3-max - temperature: 0.6 +# temperature: 0.6 + max_rps: 9 + rps_window: 10 embedding_model: default: @@ -31,6 +35,7 @@ embedding_model: vector_store: default: backend: chroma +# backend: local embedding_model: default collection_name: reme diff --git a/reme_ai/core/llm/base_llm.py b/reme_ai/core/llm/base_llm.py index 04370cc3..a741a66b 100644 --- a/reme_ai/core/llm/base_llm.py +++ b/reme_ai/core/llm/base_llm.py @@ -4,6 +4,7 @@ import asyncio import json import time from abc import ABC, abstractmethod +from collections import deque from typing import Callable, Generator, AsyncGenerator, Any from loguru import logger @@ -17,12 +18,90 @@ from ..schema import ToolCall class BaseLLM(ABC): """Abstract base class defining the standard interface for LLM interactions.""" - def __init__(self, model_name: str, max_retries: int = 3, raise_exception: bool = False, **kwargs): - """Initialize the LLM client with model configurations and retry policies.""" + def __init__(self, model_name: str, max_retries: int = 10, raise_exception: bool = False, max_rps: int | None = None, rps_window: float = 1.0, **kwargs): + """Initialize the LLM client with model configurations and retry policies. + + Args: + model_name: The name of the model to use + max_retries: Maximum number of retry attempts on failure + raise_exception: Whether to raise exceptions or return default values + max_rps: Maximum requests allowed within the time window. If None, no rate limiting is applied. + rps_window: Time window in seconds for rate limiting (default: 1.0). + For example: max_rps=10, rps_window=5.0 means max 10 requests in 5 seconds. + **kwargs: Additional model-specific parameters + """ self.model_name: str = model_name self.max_retries: int = max_retries self.raise_exception: bool = raise_exception + self.max_rps: int | None = max_rps + self.rps_window: float = rps_window self.kwargs: dict = kwargs + + # Rate limiting state - using deque for efficient O(1) operations + self._request_timestamps: deque = deque() + self._rate_limit_lock = asyncio.Lock() # For async rate limiting + import threading + self._rate_limit_lock_sync = threading.Lock() # For sync rate limiting + + async def _wait_for_rate_limit(self): + """Async rate limiting: wait if necessary to respect max_rps constraint within the time window.""" + if self.max_rps is None: + return + + async with self._rate_limit_lock: + current_time = time.time() + + # Remove timestamps older than the time window + while self._request_timestamps and current_time - self._request_timestamps[0] >= self.rps_window: + self._request_timestamps.popleft() + + # If we've reached the rate limit, wait until we can proceed + if len(self._request_timestamps) >= self.max_rps: + # Calculate how long to wait + oldest_timestamp = self._request_timestamps[0] + wait_time = self.rps_window - (current_time - oldest_timestamp) + + if wait_time > 0: + logger.debug(f"Rate limit reached ({self.max_rps} requests in {self.rps_window}s). Waiting {wait_time:.3f}s") + await asyncio.sleep(wait_time) + + # Clean up old timestamps after waiting + current_time = time.time() + while self._request_timestamps and current_time - self._request_timestamps[0] >= self.rps_window: + self._request_timestamps.popleft() + + # Record this request + self._request_timestamps.append(time.time()) + + def _wait_for_rate_limit_sync(self): + """Synchronous rate limiting: wait if necessary to respect max_rps constraint within the time window.""" + if self.max_rps is None: + return + + with self._rate_limit_lock_sync: + current_time = time.time() + + # Remove timestamps older than the time window + while self._request_timestamps and current_time - self._request_timestamps[0] >= self.rps_window: + self._request_timestamps.popleft() + + # If we've reached the rate limit, wait until we can proceed + if len(self._request_timestamps) >= self.max_rps: + # Calculate how long to wait + oldest_timestamp = self._request_timestamps[0] + wait_time = self.rps_window - (current_time - oldest_timestamp) + + if wait_time > 0: + logger.debug(f"Rate limit reached ({self.max_rps} requests in {self.rps_window}s). Waiting {wait_time:.3f}s") + time.sleep(wait_time) + + # Clean up old timestamps after waiting + current_time = time.time() + while self._request_timestamps and current_time - self._request_timestamps[0] >= self.rps_window: + self._request_timestamps.popleft() + + # Record this request + self._request_timestamps.append(time.time()) @staticmethod def _accumulate_tool_call_chunk(tool_call, ret_tools: list[ToolCall]): @@ -68,9 +147,18 @@ class BaseLLM(ABC): messages: list[Message], tools: list[ToolCall] | None = None, log_params: bool = True, + model_name: str | None = None, **kwargs, ) -> dict: - """Construct provider-specific parameters for streaming API requests.""" + """Construct provider-specific parameters for streaming API requests. + + Args: + messages: List of conversation messages + tools: Optional list of tool calls + log_params: Whether to log parameters + model_name: Optional model name to override self.model_name + **kwargs: Additional parameters + """ async def _stream_chat( self, @@ -94,10 +182,21 @@ class BaseLLM(ABC): self, messages: list[Message], tools: list[ToolCall] | None = None, + model_name: str | None = None, **kwargs, ) -> AsyncGenerator[StreamChunk, None]: - """Public async interface for streaming chat completions with retries.""" - stream_kwargs = self._build_stream_kwargs(messages, tools, **kwargs) + """Public async interface for streaming chat completions with retries. + + Args: + messages: List of conversation messages + tools: Optional list of tool calls + model_name: Optional model name to override self.model_name + **kwargs: Additional parameters + """ + # Apply rate limiting before making the request + await self._wait_for_rate_limit() + + stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs) for i in range(self.max_retries): try: @@ -121,10 +220,21 @@ class BaseLLM(ABC): self, messages: list[Message], tools: list[ToolCall] | None = None, + model_name: str | None = None, **kwargs, ) -> Generator[StreamChunk, None, None]: - """Public synchronous interface for streaming chat completions with retries.""" - stream_kwargs = self._build_stream_kwargs(messages, tools, **kwargs) + """Public synchronous interface for streaming chat completions with retries. + + Args: + messages: List of conversation messages + tools: Optional list of tool calls + model_name: Optional model name to override self.model_name + **kwargs: Additional parameters + """ + # Apply rate limiting before making the request + self._wait_for_rate_limit_sync() + + stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs) for i in range(self.max_retries): try: @@ -148,9 +258,18 @@ class BaseLLM(ABC): messages: list[Message], tools: list[ToolCall] | None = None, enable_stream_print: bool = False, + model_name: str | None = None, **kwargs, ) -> Message: - """Internal async method to aggregate a full response by consuming the stream.""" + """Internal async method to aggregate a full response by consuming the stream. + + Args: + messages: List of conversation messages + tools: Optional list of tool calls + enable_stream_print: Whether to print stream chunks + model_name: Optional model name to override self.model_name + **kwargs: Additional parameters + """ state = { "enter_think": False, "enter_answer": False, @@ -159,7 +278,7 @@ class BaseLLM(ABC): "tool_calls": [], } - stream_kwargs = self._build_stream_kwargs(messages, tools, **kwargs) + stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs) async for stream_chunk in self._stream_chat(messages=messages, tools=tools, stream_kwargs=stream_kwargs): # Process stream chunk if stream_chunk.chunk_type is ChunkEnum.USAGE: @@ -207,9 +326,18 @@ class BaseLLM(ABC): messages: list[Message], tools: list[ToolCall] | None = None, enable_stream_print: bool = False, + model_name: str | None = None, **kwargs, ) -> Message: - """Internal synchronous method to aggregate a full response by consuming the stream.""" + """Internal synchronous method to aggregate a full response by consuming the stream. + + Args: + messages: List of conversation messages + tools: Optional list of tool calls + enable_stream_print: Whether to print stream chunks + model_name: Optional model name to override self.model_name + **kwargs: Additional parameters + """ state = { "enter_think": False, "enter_answer": False, @@ -218,7 +346,7 @@ class BaseLLM(ABC): "tool_calls": [], } - stream_kwargs = self._build_stream_kwargs(messages, tools, **kwargs) + stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs) for stream_chunk in self._stream_chat_sync(messages=messages, tools=tools, stream_kwargs=stream_kwargs): # Process stream chunk if stream_chunk.chunk_type is ChunkEnum.USAGE: @@ -268,21 +396,66 @@ class BaseLLM(ABC): enable_stream_print: bool = False, callback_fn: Callable[[Message], Any] | None = None, default_value: Any = None, + model_name: str | None = None, **kwargs, ) -> Message | Any: - """Perform an async chat completion with integrated retries and error handling.""" + """Perform an async chat completion with integrated retries and error handling. + + Args: + messages: List of conversation messages + tools: Optional list of tool calls + enable_stream_print: Whether to print stream chunks + callback_fn: Optional callback function to process the result + default_value: Default value to return on error + model_name: Optional model name to override self.model_name + **kwargs: Additional parameters + """ + # Use the provided model_name or fall back to self.model_name + effective_model = model_name if model_name is not None else self.model_name + for i in range(self.max_retries): try: + # Apply rate limiting before making the request + await self._wait_for_rate_limit() + result = await self._chat( messages=messages, tools=tools, enable_stream_print=enable_stream_print, + model_name=model_name, **kwargs, ) return callback_fn(result) if callback_fn else result except Exception as e: - logger.exception(f"chat with model={self.model_name} encounter error with e={e.args}") + # Check if this is an inappropriate content error + error_message = str(e.args[0]) if e.args else str(e) + is_inappropriate_content = "inappropriate content" in error_message.lower() + is_rate_limit_error = "request rate increased too quickly" in error_message.lower() + + if is_inappropriate_content: + logger.error(f"chat with model={effective_model} detected inappropriate content error") + logger.error("=" * 80) + logger.error("Full message content that triggered the error:") + logger.error("=" * 80) + for idx, msg in enumerate(messages): + logger.error(f"Message {idx + 1} [role={msg.role}]:") + logger.error(f"Content: {msg.content}") + if msg.reasoning_content: + logger.error(f"Reasoning: {msg.reasoning_content}") + if msg.tool_calls: + logger.error(f"Tool calls: {msg.tool_calls}") + logger.error("-" * 80) + logger.error("=" * 80) + # Return empty Message immediately without retrying + return Message(role=Role.ASSISTANT, content="") + + if is_rate_limit_error: + logger.warning(f"chat with model={effective_model} hit rate limit, sleeping for 60s before retry (attempt {i + 1}/{self.max_retries})") + await asyncio.sleep(60) + continue + + logger.exception(f"chat with model={effective_model} encounter error with e={e.args}") if i == self.max_retries - 1: if self.raise_exception: @@ -299,21 +472,66 @@ class BaseLLM(ABC): enable_stream_print: bool = False, callback_fn: Callable[[Message], Any] | None = None, default_value: Any = None, + model_name: str | None = None, **kwargs, ) -> Message | Any: - """Perform a synchronous chat completion with integrated retries and error handling.""" + """Perform a synchronous chat completion with integrated retries and error handling. + + Args: + messages: List of conversation messages + tools: Optional list of tool calls + enable_stream_print: Whether to print stream chunks + callback_fn: Optional callback function to process the result + default_value: Default value to return on error + model_name: Optional model name to override self.model_name + **kwargs: Additional parameters + """ + # Use the provided model_name or fall back to self.model_name + effective_model = model_name if model_name is not None else self.model_name + for i in range(self.max_retries): try: + # Apply rate limiting before making the request + self._wait_for_rate_limit_sync() + result = self._chat_sync( messages=messages, tools=tools, enable_stream_print=enable_stream_print, + model_name=model_name, **kwargs, ) return callback_fn(result) if callback_fn else result except Exception as e: - logger.exception(f"chat sync with model={self.model_name} encounter error with e={e.args}") + # Check if this is an inappropriate content error + error_message = str(e.args[0]) if e.args else str(e) + is_inappropriate_content = "inappropriate content" in error_message.lower() + is_rate_limit_error = "request rate increased too quickly" in error_message.lower() + + if is_inappropriate_content: + logger.error(f"chat sync with model={effective_model} detected inappropriate content error") + logger.error("=" * 80) + logger.error("Full message content that triggered the error:") + logger.error("=" * 80) + for idx, msg in enumerate(messages): + logger.error(f"Message {idx + 1} [role={msg.role}]:") + logger.error(f"Content: {msg.content}") + if msg.reasoning_content: + logger.error(f"Reasoning: {msg.reasoning_content}") + if msg.tool_calls: + logger.error(f"Tool calls: {msg.tool_calls}") + logger.error("-" * 80) + logger.error("=" * 80) + # Return empty Message immediately without retrying + return Message(role=Role.ASSISTANT, content="") + + if is_rate_limit_error: + logger.warning(f"chat sync with model={effective_model} hit rate limit, sleeping for 60s before retry (attempt {i + 1}/{self.max_retries})") + time.sleep(60) + continue + + logger.exception(f"chat sync with model={effective_model} encounter error with e={e.args}") if i == self.max_retries - 1: if self.raise_exception: diff --git a/reme_ai/core/llm/lite_llm.py b/reme_ai/core/llm/lite_llm.py index a04b91f3..88177184 100644 --- a/reme_ai/core/llm/lite_llm.py +++ b/reme_ai/core/llm/lite_llm.py @@ -36,12 +36,24 @@ class LiteLLM(BaseLLM): messages: list[Message], tools: list[ToolCall] | None = None, log_params: bool = True, + model_name: str | None = None, **kwargs, ) -> dict: - """Construct and log the parameters dictionary for LiteLLM API calls.""" + """Construct and log the parameters dictionary for LiteLLM API calls. + + Args: + messages: List of conversation messages + tools: Optional list of tool calls + log_params: Whether to log parameters + model_name: Optional model name to override self.model_name + **kwargs: Additional parameters + """ + # Use the provided model_name or fall back to self.model_name + effective_model = model_name if model_name is not None else self.model_name + # Construct the API parameters by merging multiple sources llm_kwargs = { - "model": self.model_name, + "model": effective_model, "messages": [x.simple_dump() for x in messages], "tools": [x.simple_input_dump() for x in tools] if tools else None, "stream": True, diff --git a/reme_ai/core/llm/openai_llm.py b/reme_ai/core/llm/openai_llm.py index ebae540c..9e9ffc4a 100644 --- a/reme_ai/core/llm/openai_llm.py +++ b/reme_ai/core/llm/openai_llm.py @@ -41,12 +41,24 @@ class OpenAILLM(BaseLLM): messages: list[Message], tools: list[ToolCall] | None = None, log_params: bool = True, + model_name: str | None = None, **kwargs, ) -> dict: - """Construct the parameter dictionary for the OpenAI Chat Completions API call.""" + """Construct the parameter dictionary for the OpenAI Chat Completions API call. + + Args: + messages: List of conversation messages + tools: Optional list of tool calls + log_params: Whether to log parameters + model_name: Optional model name to override self.model_name + **kwargs: Additional parameters + """ + # Use the provided model_name or fall back to self.model_name + effective_model = model_name if model_name is not None else self.model_name + # Construct the API parameters by merging multiple sources llm_kwargs = { - "model": self.model_name, + "model": effective_model, "messages": [x.simple_dump() for x in messages], "tools": [x.simple_input_dump() for x in tools] if tools else None, "stream": True, diff --git a/reme_ai/core/schema/message.py b/reme_ai/core/schema/message.py index 1748f961..6c3299e7 100644 --- a/reme_ai/core/schema/message.py +++ b/reme_ai/core/schema/message.py @@ -2,6 +2,7 @@ import datetime import json +import re from pydantic import BaseModel, ConfigDict, Field, model_validator @@ -120,6 +121,7 @@ class Message(BaseModel): use_name: bool = False, add_reasoning: bool = True, add_tools: bool = True, + strip_markdown_headers: bool = False, ) -> str: """Generates a human-readable string representation of the message.""" prefix = f"round{index} " if index is not None else "" @@ -128,17 +130,23 @@ class Message(BaseModel): lines = [f"{prefix}{time_str}{header}"] + def strip_md_func(line): + if strip_markdown_headers: + line = re.sub(r'\n##+ +', '\n', line) + return line + if add_reasoning and self.reasoning_content: lines.append(self.reasoning_content) if isinstance(self.content, str): - lines.append(self.content) + lines.append(strip_md_func(self.content)) + elif isinstance(self.content, list): for block in self.content: - text = ( - block.content if isinstance(block.content, str) else json.dumps(block.content, ensure_ascii=False) - ) - lines.append(str(text)) + text = block.content if isinstance(block.content, str) else \ + json.dumps(block.content, ensure_ascii=False) + text = str(text) + lines.append(strip_md_func(text)) if add_tools and self.tool_calls: for tc in self.tool_calls: diff --git a/reme_ai/core/schema/service_config.py b/reme_ai/core/schema/service_config.py index 240f4b1a..4c6eb543 100644 --- a/reme_ai/core/schema/service_config.py +++ b/reme_ai/core/schema/service_config.py @@ -98,6 +98,7 @@ class ServiceConfig(BaseModel): language: str = Field(default="") thread_pool_max_workers: int = Field(default=16) ray_max_workers: int = Field(default=-1) + init_logger: bool = Field(default=True) disabled_flows: List[str] = Field(default_factory=list) enabled_flows: List[str] = Field(default_factory=list) mcp_servers: Dict[str, dict] = Field(default_factory=dict, description="External MCP Server configuration") diff --git a/reme_ai/core/utils/llm_utils.py b/reme_ai/core/utils/llm_utils.py index 2362a278..28935c73 100644 --- a/reme_ai/core/utils/llm_utils.py +++ b/reme_ai/core/utils/llm_utils.py @@ -23,6 +23,7 @@ def format_messages(messages: list[Message | dict], enable_system: bool = False) use_name=True, add_reasoning=True, add_tools=True, + strip_markdown_headers=True, ), ) return "\n".join(formatted_lines) diff --git a/reme_ai/core/utils/logger_utils.py b/reme_ai/core/utils/logger_utils.py index 56478683..bb6a8fcf 100644 --- a/reme_ai/core/utils/logger_utils.py +++ b/reme_ai/core/utils/logger_utils.py @@ -16,7 +16,7 @@ def init_logger(log_dir: str = "logs", level: str = "INFO") -> None: os.makedirs(log_dir, exist_ok=True) # Generate filename based on the current timestamp - current_ts = datetime.now().strftime("%Y-%m-%d_%H-%M-%S") + current_ts = datetime.now().strftime("%Y-%m-%d_%H:%M:%S") log_filename = f"{current_ts}.log" log_filepath = os.path.join(log_dir, log_filename) diff --git a/reme_ai/core/vector_store/base_vector_store.py b/reme_ai/core/vector_store/base_vector_store.py index ad858658..a4a8ca8e 100644 --- a/reme_ai/core/vector_store/base_vector_store.py +++ b/reme_ai/core/vector_store/base_vector_store.py @@ -80,6 +80,10 @@ class BaseVectorStore(ABC): async def delete(self, vector_ids: str | list[str], **kwargs) -> None: """Remove specific vectors from the collection using their identifiers.""" + @abstractmethod + async def delete_all(self, **kwargs) -> None: + """Remove all vectors from the collection.""" + @abstractmethod async def update(self, nodes: VectorNode | list[VectorNode], **kwargs) -> None: """Update the data or metadata of existing vectors in the collection.""" diff --git a/reme_ai/core/vector_store/chroma_vector_store.py b/reme_ai/core/vector_store/chroma_vector_store.py index b639405f..7baa5eab 100644 --- a/reme_ai/core/vector_store/chroma_vector_store.py +++ b/reme_ai/core/vector_store/chroma_vector_store.py @@ -66,7 +66,7 @@ class ChromaVectorStore(BaseVectorStore): self.client = chromadb.HttpClient(host=host, port=port) else: if path is None: - path = "./chroma_db" + path = "./chroma_vector_store" logger.info(f"Initializing local ChromaDB at {path}") self.client = chromadb.PersistentClient( path=path, @@ -316,6 +316,22 @@ class ChromaVectorStore(BaseVectorStore): await self._run_sync_in_executor(_delete) logger.info(f"Deleted {len(vector_ids)} nodes from {self.collection_name}") + async def delete_all(self, **kwargs): + """Remove all vectors from the collection.""" + + def _delete_all(): + # Get all IDs in the collection + result = self.collection.get() + if result and result.get("ids"): + ids = result["ids"] + if ids: + self.collection.delete(ids=ids) + return len(ids) + return 0 + + count = await self._run_sync_in_executor(_delete_all) + logger.info(f"Deleted all {count} nodes from {self.collection_name}") + async def update(self, nodes: VectorNode | list[VectorNode], **kwargs): """Update existing vector nodes with new content or metadata.""" if isinstance(nodes, VectorNode): diff --git a/reme_ai/core/vector_store/es_vector_store.py b/reme_ai/core/vector_store/es_vector_store.py index 6f52b3cc..16226749 100644 --- a/reme_ai/core/vector_store/es_vector_store.py +++ b/reme_ai/core/vector_store/es_vector_store.py @@ -320,6 +320,24 @@ class ESVectorStore(BaseVectorStore): if refresh: await self.client.indices.refresh(index=self.collection_name) + async def delete_all(self, **kwargs): + """Remove all vectors from the collection. + + Args: + **kwargs: Additional deletion parameters. + """ + response = await self.client.delete_by_query( + index=self.collection_name, + body={"query": {"match_all": {}}}, + ) + + deleted_count = response.get("deleted", 0) + logger.info(f"Deleted all {deleted_count} documents from {self.collection_name}") + + refresh = kwargs.get("refresh", True) + if refresh: + await self.client.indices.refresh(index=self.collection_name) + async def update(self, nodes: VectorNode | list[VectorNode], refresh: bool = True, **kwargs): """Update existing documents with new content or metadata. diff --git a/reme_ai/core/vector_store/local_vector_store.py b/reme_ai/core/vector_store/local_vector_store.py index a4f6d41b..cce3cae2 100644 --- a/reme_ai/core/vector_store/local_vector_store.py +++ b/reme_ai/core/vector_store/local_vector_store.py @@ -223,6 +223,24 @@ class LocalVectorStore(BaseVectorStore): logger.info(f"Deleted {deleted_count} nodes from {self.collection_name}") + async def delete_all(self, **kwargs): + """Remove all vectors from the collection.""" + col_path = self._get_collection_path(self.collection_name) + + if not col_path.exists(): + logger.warning(f"Collection {self.collection_name} does not exist") + return + + deleted_count = 0 + for file_path in col_path.glob("*.json"): + try: + file_path.unlink() + deleted_count += 1 + except Exception as e: + logger.warning(f"Failed to delete file {file_path}: {e}") + + logger.info(f"Deleted all {deleted_count} nodes from {self.collection_name}") + async def update(self, nodes: VectorNode | list[VectorNode], **kwargs): """Update existing vector nodes with new data or embeddings.""" if isinstance(nodes, VectorNode): diff --git a/reme_ai/core/vector_store/pgvector_store.py b/reme_ai/core/vector_store/pgvector_store.py index c02a0abd..a23c84b9 100644 --- a/reme_ai/core/vector_store/pgvector_store.py +++ b/reme_ai/core/vector_store/pgvector_store.py @@ -357,6 +357,16 @@ class PGVectorStore(BaseVectorStore): logger.info(f"Deleted {len(vector_ids)} documents from {self.collection_name}") + async def delete_all(self, **kwargs): + """Remove all vectors from the collection.""" + await self._ensure_collection_exists() + + pool = await self._get_pool() + async with pool.acquire() as conn: + result = await conn.execute(f"DELETE FROM {self.collection_name}") + + logger.info(f"Deleted all documents from {self.collection_name}") + async def update(self, nodes: VectorNode | list[VectorNode], **kwargs): """Update existing vector nodes with new content, embeddings, or metadata.""" await self._ensure_collection_exists() diff --git a/reme_ai/core/vector_store/qdrant_vector_store.py b/reme_ai/core/vector_store/qdrant_vector_store.py index fa48950a..1ac4db64 100644 --- a/reme_ai/core/vector_store/qdrant_vector_store.py +++ b/reme_ai/core/vector_store/qdrant_vector_store.py @@ -331,6 +331,21 @@ class QdrantVectorStore(BaseVectorStore): logger.info(f"Deleted {len(point_ids)} documents from {self.collection_name}") + async def delete_all(self, **kwargs: Any): + """Remove all vectors from the collection.""" + wait = kwargs.get("wait", True) + + # Delete all points by using an empty filter (matches all) + from qdrant_client.models import FilterSelector + + await self.client.delete( + collection_name=self.collection_name, + points_selector=FilterSelector(filter=Filter(must=[])), + wait=wait, + ) + + logger.info(f"Deleted all documents from {self.collection_name}") + async def update(self, nodes: VectorNode | list[VectorNode], **kwargs: Any): """Update existing vector nodes with new content or metadata.""" if isinstance(nodes, VectorNode): diff --git a/reme_ai/mem_agent/base_memory_agent.py b/reme_ai/mem_agent/base_memory_agent.py index 0a65834f..f2ff5587 100644 --- a/reme_ai/mem_agent/base_memory_agent.py +++ b/reme_ai/mem_agent/base_memory_agent.py @@ -36,6 +36,8 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): self.messages: list[Message] = [] self.success: bool = True + + self.retrieved_nodes: list[MemoryNode] = [] self.memory_nodes: list[MemoryNode | str] = [] def _build_tool_call(self) -> ToolCall: @@ -130,7 +132,7 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): tool_copy.tool_call.id = tool_call.id tool_list.append(tool_copy) kwargs.update(tool_call.argument_dict) - self.submit_async_task(tool_copy.call, **kwargs) + self.submit_async_task(tool_copy.call, retrieved_nodes=self.retrieved_nodes, **kwargs) if self.tool_call_interval > 0: await asyncio.sleep(self.tool_call_interval) @@ -166,17 +168,18 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta): return messages, success async def execute(self): + for i, tool in enumerate(self.tools): + logger.info( + f"[{self.__class__.__name__}] step0.{i} " + f"tool_call={json.dumps(tool.tool_call.simple_input_dump(), ensure_ascii=False)}", + ) + messages = await self.build_messages() for i, message in enumerate(messages): logger.info( f"[{self.__class__.__name__}] step0.{i} {message.role} {message.name or ''} " f"{message.simple_dump(enable_json_dump=True)}", ) - for i, tool in enumerate(self.tools): - logger.info( - f"[{self.__class__.__name__}] step0.{i} " - f"tool_call={json.dumps(tool.tool_call.simple_input_dump(), ensure_ascii=False)}", - ) self.messages, self.success = await self.react(messages) if self.success and self.messages: diff --git a/reme_ai/mem_agent/retriever/reme_retriever.py b/reme_ai/mem_agent/retriever/reme_retriever.py index 6841432a..f3700b1e 100644 --- a/reme_ai/mem_agent/retriever/reme_retriever.py +++ b/reme_ai/mem_agent/retriever/reme_retriever.py @@ -14,7 +14,8 @@ class ReMeRetriever(BaseMemoryAgent): """Memory agent that retrieves and builds messages with meta memory context.""" def __init__(self, meta_memories: list[dict] | None = None, **kwargs): - super().__init__(**kwargs) + # super().__init__(prompt_name="", **kwargs) + super().__init__(prompt_name="reme_retriever2", **kwargs) self.meta_memories: list[dict] = meta_memories or [] async def _read_meta_memories(self) -> str: @@ -28,20 +29,38 @@ class ReMeRetriever(BaseMemoryAgent): await op.call() return str(op.output) + # async def build_messages1(self) -> List[Message]: + # """Build messages with system prompt and user message.""" + # meta_memory_info = await self._read_meta_memories() + # system_prompt = self.prompt_format( + # prompt_name="system_prompt", + # now_time=get_now_time(), + # meta_memory_info=meta_memory_info, + # ) + + # messages = [Message(role=Role.SYSTEM, content=system_prompt)] + # if self.context.get("query"): + # messages.append(Message(role=Role.USER, content=self.context.query)) + # elif self.context.get("messages"): + # messages.extend([Message(**m) for m in self.context.messages]) + # else: + # raise ValueError("input must have either `query` or `messages`") + + # return messages + async def build_messages(self) -> List[Message]: """Build messages with system prompt and user message.""" - meta_memory_info = await self._read_meta_memories() if self.context.get("query"): context = self.context.query elif self.context.get("messages"): - messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages] - context = self.description + format_messages(messages) + context = format_messages(self.context.messages) else: raise ValueError("input must have either `query` or `messages`") + system_prompt = self.prompt_format( prompt_name="system_prompt", now_time=get_now_time(), - meta_memory_info=meta_memory_info, + meta_memory_info=await self._read_meta_memories(), context=context, ) @@ -49,4 +68,5 @@ class ReMeRetriever(BaseMemoryAgent): Message(role=Role.SYSTEM, content=system_prompt), Message(role=Role.USER, content=self.get_prompt("user_message")), ] + return messages diff --git a/reme_ai/mem_agent/retriever/reme_retriever.yaml b/reme_ai/mem_agent/retriever/reme_retriever.yaml index 056b5138..8ab3a76c 100644 --- a/reme_ai/mem_agent/retriever/reme_retriever.yaml +++ b/reme_ai/mem_agent/retriever/reme_retriever.yaml @@ -6,10 +6,9 @@ tool: | semantic searches across different memory types to find the most relevant memories. system_prompt: | - You are a memory agent. Please analyze the context, retrieve relevant memories when needed, and return a summary of the retrieved memories to assist in answering the user's question. + You are a memory agent. Please analyze the context, retrieve relevant memories when needed, and directly answer the user's question based on the retrieved information. - ## Context - {context} + **CRITICAL**: You must ONLY answer based on the retrieved memories. DO NOT fabricate, infer, or add any information that is not explicitly present in the retrieved memories. If the retrieved memories do not contain enough information to answer the question, you must acknowledge this limitation. ## Current Time {now_time} @@ -34,6 +33,7 @@ system_prompt: | * Choose the optimal combination strategy based on the retrieval scenario. - **Important**: When retrieving tool-related memories (`memory_type` is "tool"), the query must use the tool’s exact name (not a description or paraphrase of the problem). - If retrieval results include a `ref_memory_id` and more details are needed—or if vector retrieval proves insufficient—use `read_history_memory` with the `ref_memory_id` as the `memory_id` parameter. + - **Important**: When using `read_history_memory` with multiple `ref_memory_ids`, ensure all IDs are unique and do not provide duplicate IDs. 3. **Iterate if necessary**: - If the initial retrieval fails, try alternative phrasings or perspectives. @@ -43,8 +43,5 @@ system_prompt: | 4. **Output** the result: - If no retrieval is needed, output ``. - - If relevant memories are found, clearly summarize the retrieved information. - - If multiple attempts still yield no relevant memory, output ``. - -user_message: | - Please analyze the context, retrieve relevant memories when needed, and return a summary of the retrieved memories to assist in answering the user's question. + - If relevant memories are found and you can answer the user's question, provide a concise, direct answer **strictly based on the retrieved memories only**. DO NOT add any information, inference, or speculation beyond what is explicitly stated in the retrieved memories. + - If after multiple retrieval attempts from various angles you still cannot find relevant information, output ``. diff --git a/reme_ai/mem_agent/retriever/reme_retriever2.yaml b/reme_ai/mem_agent/retriever/reme_retriever2.yaml new file mode 100644 index 00000000..bc139996 --- /dev/null +++ b/reme_ai/mem_agent/retriever/reme_retriever2.yaml @@ -0,0 +1,49 @@ +tool: | + Retrieve relevant memories from the memory bank to assist in answering questions. + Use this tool when you need to search for historical information, user preferences, + procedural knowledge, or any other stored memories that may help answer the current query. + The agent will analyze the context, determine what information is needed, and perform + semantic searches across different memory types to find the most relevant memories. + +system_prompt: | + You are a memory retrieval agent. Please analyze the context, retrieve relevant information from the memory bank when needed, and return a summary of the retrieved memories to assist in answering the user's question. + + ## Context + {context} + + ## Current Time + {now_time} + + ## Available Meta-Memories + Format: "- (): " + {meta_memory_info} + + ## Your Tasks + 1. **Analyze** the conversation context to determine whether retrieval is necessary: + - If the question can be answered directly from the existing context, output `` and stop. + - If additional information is required, proceed with retrieval. + - Consider which types of meta-memories from the "Available Meta-Memories" list are most relevant. + + 2. **Retrieve** relevant memories using `vector_retrieve_memory`: + - Select `memory_type` and `memory_target` from the "Available Meta-Memories" list. + - Clearly identify the needed information and construct appropriate queries. + - Design queries flexibly based on actual needs: + * Generate different queries for different `memory_type`/`memory_target` combinations. + * For the same combination, generate multiple queries with varied phrasings or angles if needed. + * Use the combination strategy that best fits the retrieval scenario. + - **Important**: When retrieving tool memories (`memory_type` is "tool"), use the actual tool name as the query (not a description or question). + - If retrieval results include a `ref_memory_id` and more detail is needed—or if vector retrieval proves insufficient—use `read_history_memory` with the `ref_memory_id` as the `memory_id` parameter. + + 3. **Iterate if necessary**: + - If the initial retrieval yields no matches, try alternative phrasings or perspectives. + - If multiple memory types exist, attempt retrieval across different types. + - Before concluding that no relevant memory exists, perform at least 2–3 additional retrieval attempts using varied phrasings or angles. + - If repeated vector retrievals still fail to provide adequate information, use `read_history_memory` to fetch the original message content. + + 4. **Output** the result: + - If no retrieval is needed, output ``. + - If relevant memories are found, clearly summarize the retrieved information. + - If multiple attempts still yield no relevant memories, output ``. + +user_message: | + Please analyze the context, retrieve relevant information from the memory bank when needed, and return a summary of the retrieved memories to assist in answering the user's question. diff --git a/reme_ai/mem_agent/retriever_v2/__init__.py b/reme_ai/mem_agent/retriever_v2/__init__.py new file mode 100644 index 00000000..726ec8fe --- /dev/null +++ b/reme_ai/mem_agent/retriever_v2/__init__.py @@ -0,0 +1,5 @@ +from .reme_retriever_v2 import ReMeRetrieverV2 + +__all__ = [ + "ReMeRetrieverV2", +] \ No newline at end of file diff --git a/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.py b/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.py new file mode 100644 index 00000000..ee5c0aae --- /dev/null +++ b/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.py @@ -0,0 +1,67 @@ +"""ReMe retriever v2 that autonomously retrieves memories from multiple angles.""" + +from typing import List + +from ..base_memory_agent import BaseMemoryAgent +from ...core.context import C +from ...core.enumeration import Role +from ...core.schema import Message +from ...core.utils import format_messages + + +@C.register_op() +class ReMeRetrieverV2(BaseMemoryAgent): + """Memory agent that autonomously retrieves memories from multiple angles. + + This retriever: + - Directly queries memories based on user questions without time constraints + - Tries multiple retrieval strategies: direct vector search, metadata filtering, partial filtering + - Attempts at least 3 vector retrievals from different perspectives + - Falls back to read_history if vector retrieval doesn't find sufficient information + """ + + def __init__(self, meta_memories: list[dict] | None = None, **kwargs): + # Check if ReadHistory tool is available in the tools list + tools = kwargs.get('tools', []) + has_read_history = any(tool.__class__.__name__ == 'ReadHistory' for tool in tools) + + # Use simple prompt if ReadHistory is not available + if not has_read_history: + super().__init__(prompt_name="reme_retriever_v2_simple", **kwargs) + else: + super().__init__(**kwargs) + + self.meta_memories: list[dict] = meta_memories or [] + + async def _read_meta_memories(self) -> str: + """Fetch all meta-memory entries that define specialized memory agents.""" + from ...mem_tool import ReadMetaMemory + + op = ReadMetaMemory(enable_identity_memory=False) + if self.meta_memories: + return op.format_memory_metadata(self.meta_memories) + else: + await op.call() + return str(op.output) + + async def build_messages(self) -> List[Message]: + """Build messages with system prompt and user message.""" + if self.context.get("query"): + context = self.context.query + elif self.context.get("messages"): + context = format_messages(self.context.messages) + else: + raise ValueError("input must have either `query` or `messages`") + + system_prompt = self.prompt_format( + prompt_name="system_prompt", + meta_memory_info=await self._read_meta_memories(), + context=context, + ) + + messages = [ + Message(role=Role.SYSTEM, content=system_prompt), + Message(role=Role.USER, content=self.get_prompt("user_message")), + ] + + return messages diff --git a/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.yaml b/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.yaml new file mode 100644 index 00000000..281796d6 --- /dev/null +++ b/reme_ai/mem_agent/retriever_v2/reme_retriever_v2.yaml @@ -0,0 +1,125 @@ +tool: | + Autonomously retrieve relevant memories from multiple angles to answer user questions. + This retriever will: + - Try multiple vector search strategies (direct, metadata-filtered, partial) + - Attempt at least 3 different retrieval approaches before giving up + - Fall back to reading original conversation history if vector search is insufficient + - Clearly state "I don't know" if information cannot be found after exhaustive searching + - NEVER hallucinate or fabricate information not present in retrieved memories + Use this when you need comprehensive memory retrieval with persistent searching. + +system_prompt: | + You are an autonomous memory retrieval agent. Your task is to persistently search for relevant memories from multiple angles to answer the user's question. + + ## Available Meta Memories + Format: "- (): " + {meta_memory_info} + + ## User Context + {context} + + ## Your Retrieval Strategy + + You MUST use the `retrieve_memories` tool to search for relevant information. This is a MANDATORY step - do not skip it. + + 1. **Multi-Angle Vector Retrieval** (REQUIRED - at least 3 attempts): + You must try AT LEAST 3 different retrieval approaches using `retrieve_memories`: + + a) **Direct Vector Search**: Use the user's question directly or with minimal reformulation + - Query the most relevant memory_type and memory_target + - Use straightforward query phrasing + + b) **Alternative Phrasing**: Reformulate the query from a different angle + - Use synonyms or different expressions + - Break down complex questions into simpler components + - Try more specific or more general queries + + c) **Metadata-Filtered Search**: Add metadata filters to narrow down results + - **Time-based filtering**: Use year/month/day metadata fields to filter by time periods + * Example: {{"year": 2024}} for memories from 2024 + * Example: {{"year": 2024, "month": 5}} for memories from May 2024 + * Example: {{"year": 2024, "month": 5, "day": 15}} for memories from a specific date + - Combine vector search with metadata constraints + - Try partial metadata filtering if full filtering yields nothing (e.g., only year, or year+month) + + d) **Cross-Memory-Type Search**: If applicable, search across different memory types + - Try different memory_type and memory_target combinations + - Some information might be stored in unexpected memory categories + + e) **Keyword Extraction**: Extract key entities/concepts and search for them + - Identify important names, places, concepts + - Search for each key element separately + + 2. **Evaluate Retrieval Results** (After each attempt): + - Review what memories were returned + - Assess if they contain sufficient information to answer the question + - If insufficient, identify what's missing and adjust your next query accordingly + - Track which retrieval strategies you've already tried + + 3. **Persist Through Failures**: + - DO NOT give up after 1-2 failed attempts + - If a retrieval returns no results or irrelevant results, try a different approach + - Consider that the information might be phrased differently than expected + - Be creative with query reformulation + + 4. **Fallback to History Reading** (Only after 3+ vector retrieval attempts): + - If after at least 3 different vector retrieval attempts you still lack sufficient information: + * If any retrieved memories contain `ref_memory_id`, use `read_history` to read the original conversation + * Use `read_history` with the `ref_memory_id` to get complete context + * This can reveal details that weren't captured in the memory summaries + + 5. **Answer the Question**: + - Once you have sufficient information, provide a direct answer based ONLY on retrieved memories + - DO NOT fabricate, guess, or infer information not present in the memories + - **CRITICAL**: If after 3+ retrieval attempts you still cannot find relevant information: + * Simply state: "I don't know. After searching from multiple angles, I could not find relevant information to answer this question." + * DO NOT make up answers or hallucinate information + * DO NOT provide speculative or guessed responses + * It is better to say "I don't know" than to provide incorrect information + + ## Important Guidelines + + - **Be Persistent**: Always try at least 3 different retrieval strategies before concluding no information exists + - **Be Creative**: If one query approach fails, think of alternative ways to phrase or decompose the question + - **Use Tools**: You MUST use `retrieve_memories` for vector search. Use `read_history` if you have `ref_memory_id` and need more details + - **No Hallucination**: NEVER fabricate, guess, or hallucinate information. Only answer based on what you actually retrieved from memories + - **Admit When You Don't Know**: If after 3+ attempts you cannot find relevant information, clearly say "I don't know" rather than making up an answer + - **Track Your Attempts**: Keep count of how many different retrieval strategies you've tried + - **Metadata Awareness**: Utilize metadata filters when they might help narrow down results + * Memories store time information in metadata as year/month/day fields + * Use time-based filters when the question involves specific time periods or dates + * Try progressive filtering: start with year, then add month, then day if needed + + ## Example Retrieval Flow + + **Example 1: Simple Query** + Attempt 1: Direct query "user's favorite food" + → Result: No relevant memories found + + Attempt 2: Reformulated query "what does user like to eat" + → Result: Some memories about meals, but not specific preferences + + Attempt 3: Keyword search "food preferences" with metadata filter + → Result: Found relevant memory with ref_memory_id + + Attempt 4: Use read_history with ref_memory_id to get full context + → Result: Found detailed conversation about favorite foods + + Answer: [Provide answer based on retrieved information] + + **Example 2: Time-based Query** + Question: "What did the user do last summer?" + + Attempt 1: Direct query "user activities summer" with metadata {{"year": 2025, "month": [6, 7, 8]}} + → Result: Found some vacation memories + + Attempt 2: Broader query "user summer vacation travel" with metadata {{"year": 2025}} + → Result: Found additional travel-related memories + + Attempt 3: Use read_history for memories with ref_memory_id to get detailed context + → Result: Complete picture of summer activities + + Answer: [Provide answer based on retrieved information] + +user_message: | + Please retrieve relevant memories and answer the question. Remember to try multiple retrieval approaches before giving up. diff --git a/reme_ai/mem_agent/retriever_v2/reme_retriever_v2_simple.yaml b/reme_ai/mem_agent/retriever_v2/reme_retriever_v2_simple.yaml new file mode 100644 index 00000000..f8a4b7f5 --- /dev/null +++ b/reme_ai/mem_agent/retriever_v2/reme_retriever_v2_simple.yaml @@ -0,0 +1,115 @@ +tool: | + Autonomously retrieve relevant memories from multiple angles to answer user questions. + This retriever will: + - Try multiple vector search strategies (direct, metadata-filtered, partial) + - Attempt at least 3 different retrieval approaches before giving up + - Clearly state "I don't know" if information cannot be found after exhaustive searching + - NEVER hallucinate or fabricate information not present in retrieved memories + Use this when you need comprehensive memory retrieval with persistent searching. + +system_prompt: | + You are an autonomous memory retrieval agent. Your task is to persistently search for relevant memories from multiple angles to answer the user's question. + + ## Available Meta Memories + Format: "- (): " + {meta_memory_info} + + ## User Context + {context} + + ## Your Retrieval Strategy + + You MUST use the `retrieve_memories` tool to search for relevant information. This is a MANDATORY step - do not skip it. + + 1. **Multi-Angle Vector Retrieval** (REQUIRED - at least 3 attempts): + You must try AT LEAST 3 different retrieval approaches using `retrieve_memories`: + + a) **Direct Vector Search**: Use the user's question directly or with minimal reformulation + - Query the most relevant memory_type and memory_target + - Use straightforward query phrasing + + b) **Alternative Phrasing**: Reformulate the query from a different angle + - Use synonyms or different expressions + - Break down complex questions into simpler components + - Try more specific or more general queries + + c) **Metadata-Filtered Search**: Add metadata filters to narrow down results + - **Time-based filtering**: Use year/month/day metadata fields to filter by time periods + * Example: {{"year": 2024}} for memories from 2024 + * Example: {{"year": 2024, "month": 5}} for memories from May 2024 + * Example: {{"year": 2024, "month": 5, "day": 15}} for memories from a specific date + - Combine vector search with metadata constraints + - Try partial metadata filtering if full filtering yields nothing (e.g., only year, or year+month) + + d) **Cross-Memory-Type Search**: If applicable, search across different memory types + - Try different memory_type and memory_target combinations + - Some information might be stored in unexpected memory categories + + e) **Keyword Extraction**: Extract key entities/concepts and search for them + - Identify important names, places, concepts + - Search for each key element separately + + 2. **Evaluate Retrieval Results** (After each attempt): + - Review what memories were returned + - Assess if they contain sufficient information to answer the question + - If insufficient, identify what's missing and adjust your next query accordingly + - Track which retrieval strategies you've already tried + + 3. **Persist Through Failures**: + - DO NOT give up after 1-2 failed attempts + - If a retrieval returns no results or irrelevant results, try a different approach + - Consider that the information might be phrased differently than expected + - Be creative with query reformulation + + 4. **Answer the Question**: + - Once you have sufficient information, provide a direct answer based ONLY on retrieved memories + - DO NOT fabricate, guess, or infer information not present in the memories + - **CRITICAL**: If after 3+ retrieval attempts you still cannot find relevant information: + * Simply state: "I don't know. After searching from multiple angles, I could not find relevant information to answer this question." + * DO NOT make up answers or hallucinate information + * DO NOT provide speculative or guessed responses + * It is better to say "I don't know" than to provide incorrect information + + ## Important Guidelines + + - **Be Persistent**: Always try at least 3 different retrieval strategies before concluding no information exists + - **Be Creative**: If one query approach fails, think of alternative ways to phrase or decompose the question + - **Use Tools**: You MUST use `retrieve_memories` for vector search + - **No Hallucination**: NEVER fabricate, guess, or hallucinate information. Only answer based on what you actually retrieved from memories + - **Admit When You Don't Know**: If after 3+ attempts you cannot find relevant information, clearly say "I don't know" rather than making up an answer + - **Track Your Attempts**: Keep count of how many different retrieval strategies you've tried + - **Metadata Awareness**: Utilize metadata filters when they might help narrow down results + * Memories store time information in metadata as year/month/day fields + * Use time-based filters when the question involves specific time periods or dates + * Try progressive filtering: start with year, then add month, then day if needed + + ## Example Retrieval Flow + + **Example 1: Simple Query** + Attempt 1: Direct query "user's favorite food" + → Result: No relevant memories found + + Attempt 2: Reformulated query "what does user like to eat" + → Result: Some memories about meals, but not specific preferences + + Attempt 3: Keyword search "food preferences" with metadata filter + → Result: Found relevant memory + + Answer: [Provide answer based on retrieved information] + + **Example 2: Time-based Query** + Question: "What did the user do last summer?" + + Attempt 1: Direct query "user activities summer" with metadata {{"year": 2025, "month": [6, 7, 8]}} + → Result: Found some vacation memories + + Attempt 2: Broader query "user summer vacation travel" with metadata {{"year": 2025}} + → Result: Found additional travel-related memories + + Attempt 3: More specific queries about specific activities + → Result: Complete picture of summer activities + + Answer: [Provide answer based on retrieved information] + +user_message: | + Please retrieve relevant memories and answer the question. Remember to try multiple retrieval approaches before giving up. diff --git a/reme_ai/mem_agent/summarizer/personal_summarizer.py b/reme_ai/mem_agent/summarizer/personal_summarizer.py index 696ff256..352fe1ee 100644 --- a/reme_ai/mem_agent/summarizer/personal_summarizer.py +++ b/reme_ai/mem_agent/summarizer/personal_summarizer.py @@ -11,6 +11,10 @@ from ...core.utils import get_now_time, format_messages class PersonalSummarizer(BaseMemoryAgent): """Extracts and stores personal information about individuals from conversations.""" + def __init__(self, recent_top_k: int = 20, **kwargs): + super().__init__(**kwargs) + self.recent_top_k: int = recent_top_k + memory_type: MemoryType = MemoryType.PERSONAL def _build_tool_call(self) -> ToolCall: @@ -43,11 +47,22 @@ class PersonalSummarizer(BaseMemoryAgent): }, ) + async def _retrieve_recent_memories(self) -> str: + """Retrieve recent memories sorted by time_modified.""" + from ...mem_tool import RetrieveRecentMemory + + op = RetrieveRecentMemory(top_k=self.recent_top_k) + await op.call(memory_type="personal", memory_target=self.memory_target, retrieved_nodes=self.retrieved_nodes) + return op.output + async def build_messages(self) -> list[Message]: """Construct messages with context, memory_target, and memory_type information.""" + await self._retrieve_recent_memories() + system_prompt = self.prompt_format( prompt_name="system_prompt", now_time=get_now_time(), + recent_memories="\n".join([n.format_memory() for n in self.retrieved_nodes]), context=self.description + "\n" + format_messages(self.get_messages()), memory_type=self.memory_type.value, memory_target=self.memory_target, diff --git a/reme_ai/mem_agent/summarizer/personal_summarizer.yaml b/reme_ai/mem_agent/summarizer/personal_summarizer.yaml index 7ed0b091..cdad37f3 100644 --- a/reme_ai/mem_agent/summarizer/personal_summarizer.yaml +++ b/reme_ai/mem_agent/summarizer/personal_summarizer.yaml @@ -8,12 +8,17 @@ tool: | 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. + **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: {context} ## Current Time: {now_time} + ## Recent Memories: + {recent_memories} + ## 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. @@ -22,22 +27,36 @@ system_prompt: | 1. **Analyze and Extract** potential memories from the dialogue 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. - If the dialogue is casual chatter or contains no valuable information, output `` and stop. - - Extract key information using clear and concise phrasing. + - Extract key information 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. - Each memory entry must be self-contained and understandable without additional context. - Avoid storing trivial or temporary information. - - Before proceeding, list all extracted memories in your response. + - **CRITICAL**: After extraction, immediately deduplicate within the extracted memories themselves - if multiple extracted items convey the same core information (even with slightly different wording), keep ONLY the most complete and accurate one. + - Before proceeding, list all deduplicated extracted memories in your response. 2. **Retrieve similar historical memories** using `vector_retrieve_memory`: - - Perform a semantic similarity search based on the extracted memories to find existing, potentially relevant memories. - - Retrieve related memories for comparison to check for duplication or associations. + - For EACH extracted memory, perform a semantic similarity search to find existing, potentially relevant memories (e.g., for "Person A was born on date X", search for "Person A birth date age"). + - Retrieve all related memories for thorough comparison to prevent any duplication. - 3. **Compare and Decide** on memory operations: - - Compare the newly extracted memories with historical ones to ensure the final memory store contains no duplicates or contradictions. + 3. **Compare and Decide** on memory operations with STRICT deduplication: + - Compare the newly extracted memories with **both Recent Memories and historical memories** retrieved in the previous step. + - **CRITICAL DEDUPLICATION CHECK**: Before adding ANY new memory: + - Check if the SAME INFORMATION already exists in Recent Memories or retrieved historical 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 + - **Use `update_memory` and `delete_memory` to actively deduplicate and resolve conflicts:** + - If multiple existing memories contain duplicate or overlapping information: use `delete_memory` to remove redundant ones, then use `update_memory` to consolidate all information into a single, comprehensive memory. + - If memories conflict (contradictory information): use `delete_memory` to remove outdated/incorrect ones, then use `update_memory` or `add_memory` to store the correct version. - Choose the appropriate operation based on the situation: - - If the information already exists and is consistent: skip—no action needed. - - If existing memory needs supplementation or correction: use `update_memory` to update it. - - If existing memory is outdated or incorrect: use `delete_memory` to remove it. - - If the information is entirely new: use `add_memory` to add it to the memory store. + - **If the information already exists in Recent Memories or historical memories and is consistent: SKIP—no action needed. Do NOT add duplicate memories.** + - If existing memory (recent or historical) needs supplementation with NEW details: use `update_memory` to enhance and consolidate it. + - If existing memory (recent or historical) is outdated or contradicted by new information: use `delete_memory` to remove it, then `add_memory` for the corrected version. + - If multiple memories contain similar/overlapping information: use `delete_memory` to remove duplicates, then `update_memory` to merge into one. + - If the information is entirely new and not present in either Recent Memories or historical memories: use `add_memory` to add it to the memory store. + - **When in doubt, prefer updating or consolidating existing memories over adding new ones to avoid redundancy.** 4. **Output** the result: - If no memory operation is required, output ``. @@ -46,9 +65,18 @@ system_prompt: | ## Guidelines: - Be selective: store only truly important information. - Stay concise: each memory should be clear and atomic. - - Be accurate: ensure extracted content faithfully reflects the original context. - - Avoid redundancy: always check for similar existing memories before adding new ones. + - **Be strictly accurate**: ensure extracted content faithfully reflects ONLY what is explicitly stated in the original context. DO NOT infer, extrapolate, or fabricate any details. + - **AVOID REDUNDANCY AT ALL COSTS**: This is your TOP PRIORITY. Always perform thorough deduplication: + * First, deduplicate within newly extracted memories + * Then, check against Recent Memories (provided above) + * Finally, use `vector_retrieve_memory` to check against historical memories + * **Actively use `delete_memory` to remove duplicate or conflicting memories** + * **Use `update_memory` to consolidate and integrate information from multiple memories into one** + * If information semantically matches existing memories, DO NOT add it again + * When uncertain, prefer to skip or update existing memories rather than create duplicates - Include relevant metadata (e.g., timestamps) when appropriate. + - **No assumptions**: Only store information that is directly and clearly stated in the conversation. + - **Quality over quantity**: It's better to have fewer, well-maintained memories than many duplicate ones. user_message: | Please analyze the context to determine whether important information should be extracted and stored as memory, and perform memory addition, deletion, or update operations when necessary. diff --git a/reme_ai/mem_agent/summarizer/reme_summarizer.py b/reme_ai/mem_agent/summarizer/reme_summarizer.py index 168b73ec..9eb06d84 100644 --- a/reme_ai/mem_agent/summarizer/reme_summarizer.py +++ b/reme_ai/mem_agent/summarizer/reme_summarizer.py @@ -1,12 +1,10 @@ """Orchestrator for complete memory summarization workflow across all memory types.""" -from typing import List - from loguru import logger from ..base_memory_agent import BaseMemoryAgent from ...core.context import C -from ...core.enumeration import Role +from ...core.enumeration import Role, MemoryType from ...core.schema import Message, MemoryNode, ToolCall from ...core.utils import get_now_time, format_messages @@ -20,11 +18,23 @@ class ReMeSummarizer(BaseMemoryAgent): super().__init__(**kwargs) self.enable_identity_memory = enable_identity_memory self.meta_memories: list[dict] = meta_memories or [] + + # Check if AddMetaMemory is in tools + self.enable_add_meta_memory = self._check_add_meta_memory_in_tools() + def _check_add_meta_memory_in_tools(self) -> bool: + """Check if AddMetaMemory tool is present in the tools list.""" + from ...mem_tool import AddMetaMemory + + for tool in self.tools: + if isinstance(tool, AddMetaMemory): + return True + return False + def _build_tool_call(self) -> ToolCall: return ToolCall( **{ - "description": self.get_prompt("tool"), + "description": self.prompt_format("tool", enable_add_meta_memory=self.enable_add_meta_memory), "parameters": { "type": "object", "properties": { @@ -51,22 +61,16 @@ class ReMeSummarizer(BaseMemoryAgent): }, ) - async def _add_history_memory(self) -> MemoryNode: - """Store conversation history and return the memory node.""" - from ...mem_tool import AddHistoryMemory - - op = AddHistoryMemory() - await op.call(messages=self.get_messages()) - return op.memory_nodes[0] - - @staticmethod - async def _read_identity_memory() -> str: + async def _read_identity_memory(self) -> str: """Retrieve agent's self-perception memory.""" - from ...mem_tool import ReadIdentityMemory + if self.enable_identity_memory: + from ...mem_tool import ReadIdentityMemory - op = ReadIdentityMemory() - await op.call() - return op.output + op = ReadIdentityMemory() + await op.call() + return op.output + else: + return "" async def _read_meta_memories(self) -> str: """Fetch all meta-memory entries that define specialized memory agents.""" @@ -79,29 +83,27 @@ class ReMeSummarizer(BaseMemoryAgent): await op.call() return str(op.output) - async def build_messages(self) -> List[Message]: + async def build_messages(self) -> list[Message]: """Construct initial messages with context, identity, and meta-memory information.""" - memory_node: MemoryNode = await self._add_history_memory() - self.context["ref_memory_id"] = memory_node.memory_id + messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages] + self.context["messages_formated"] = self.description + "\n" + format_messages(messages) + self.context["ref_memory_id"] = MemoryNode( + memory_type=MemoryType.HISTORY, + content=self.context["messages_formated"], + ).memory_id now_time = get_now_time() identity_memory = await self._read_identity_memory() meta_memory_info = await self._read_meta_memories() - context = self.description + "\n" + format_messages(self.get_messages()) - logger.info( - f"now_time={now_time} " - f"memory_node={memory_node.content[:100]}... " - f"identity_memory={identity_memory} " - f"meta_memory_info={meta_memory_info} " - f"context={context[:100]}", - ) + logger.info(f"now_time={now_time} identity_memory={identity_memory} meta_memory_info={meta_memory_info}") system_prompt = self.prompt_format( prompt_name="system_prompt", now_time=now_time, identity_memory=identity_memory, meta_memory_info=meta_memory_info, - context=context, + context=self.context["messages_formated"], + enable_add_meta_memory=self.enable_add_meta_memory, ) user_message = self.get_prompt("user_message") @@ -118,16 +120,13 @@ class ReMeSummarizer(BaseMemoryAgent): if system_messages: system_message = system_messages[0] - now_time = get_now_time() - identity_memory = await self._read_identity_memory() - meta_memory_info = await self._read_meta_memories() - context = self.description + "\n" + format_messages(self.get_messages()) system_message.content = self.prompt_format( prompt_name="system_prompt", - now_time=now_time, - identity_memory=identity_memory, - meta_memory_info=meta_memory_info, - context=context, + now_time=get_now_time(), + identity_memory=await self._read_identity_memory(), + meta_memory_info=await self._read_meta_memories(), + context=self.context["messages_formated"], + enable_add_meta_memory=self.enable_add_meta_memory, ) return await super()._reasoning_step(messages, step, **kwargs) @@ -140,6 +139,7 @@ class ReMeSummarizer(BaseMemoryAgent): messages=self.context.get("messages", []), description=self.context.get("description"), ref_memory_id=self.context["ref_memory_id"], + messages_formated=self.context["messages_formated"], author=self.author, **kwargs, ) diff --git a/reme_ai/mem_agent/summarizer/reme_summarizer.yaml b/reme_ai/mem_agent/summarizer/reme_summarizer.yaml index ad254342..90f18685 100644 --- a/reme_ai/mem_agent/summarizer/reme_summarizer.yaml +++ b/reme_ai/mem_agent/summarizer/reme_summarizer.yaml @@ -1,9 +1,9 @@ tool: | Orchestrate the complete memory summarization workflow for the agent. This tool receives conversation context and performs necessary memory updates including: - 1. Creating new meta-memory entries if needed - 2. Adding summary memory for quick future recall - 3. Delegating to specialized memory agents for detailed memory extraction and update + [enable_add_meta_memory]- Creating new meta-memory entries if needed + - Adding summary memory for quick future recall + - Delegating to specialized memory agents for detailed memory extraction and update system_prompt: | You are a Memory Agent responsible for performing necessary updates and summaries of the main Agent's memories based on the **context**. @@ -24,27 +24,27 @@ system_prompt: | ## Your Tasks - ### 1. Create New Meta Memory (if needed) - When the context contains significant new valuable information, first check if the Main Agent's Meta Memory already contains a corresponding `()` entry: - - If the required `()` does NOT exist in the Meta Memory, use `add_meta_memory` to create a new meta memory entry. - - For personal memories: specify `memory_type="personal"` and `memory_target=`. - - For procedural memories: specify `memory_type="procedural"` and `memory_target=`. - - Each meta memory entry will instantiate a dedicated specialized Memory Agent for that dimension. - - Only create new meta memory entries when necessary; avoid duplicating existing ones. - - ### 2. Add Summary Memory (if valuable) - When the context includes information worth remembering for quick future recall: + ### 1. Add Summary Memory - Use `add_summary_memory` to store a concise summary. - The summary should capture key points, decisions, or important facts to aid later recollection of the original conversation. - ### 3. Delegate to Specialized Memory Agents (Core Task) + [enable_add_meta_memory]### 2. Create New Meta Memory (if needed) + [enable_add_meta_memory]When the context contains significant new valuable information, first check if the Main Agent's Meta Memory already contains a corresponding `()` entry: + [enable_add_meta_memory]- If the required `()` does NOT exist in the Meta Memory, use `add_meta_memory` to create a new meta memory entry. + [enable_add_meta_memory]- For personal memories: specify `memory_type="personal"` and `memory_target=`. + [enable_add_meta_memory]- For procedural memories: specify `memory_type="procedural"` and `memory_target=`. + [enable_add_meta_memory]- Each meta memory entry will instantiate a dedicated specialized Memory Agent for that dimension. + [enable_add_meta_memory]- Only create new meta memory entries when necessary; avoid duplicating existing ones. + [enable_add_meta_memory] + + ### 3. Delegate to Specialized Memory Agents You do not need to summarize or update memories yourself. Instead, analyze the context, identify which memory dimensions (memory_type + memory_target) from the existing meta memory require updates, and delegate using `hands_off`: - The parameters of `hands_off` (`memory_type` and `memory_target`) must exactly match an existing entry in the "Main Agent's Meta Memory" listed above. - You may delegate concurrently to multiple specialized agents to enable parallel memory processing. - Each specialized agent will perform detailed memory extraction, addition, updating, or deletion within its assigned dimension. ## Output Requirements - - If the context contains no memorable information (e.g., simple greetings or meaningless small talk), output ``. + - If the context contains no memorable information (e.g., simple greetings), output ``. - If any memory operations were performed, briefly summarize what was done. user_message: | diff --git a/reme_ai/mem_agent/summarizer_v2/__init__.py b/reme_ai/mem_agent/summarizer_v2/__init__.py new file mode 100644 index 00000000..4401fe74 --- /dev/null +++ b/reme_ai/mem_agent/summarizer_v2/__init__.py @@ -0,0 +1,6 @@ +"""Simplified V2 summarizers for memory management.""" + +from .reme_summarizer_v2 import ReMeSummarizerV2 +from .personal_summarizer_v2 import PersonalSummarizerV2 + +__all__ = ["ReMeSummarizerV2", "PersonalSummarizerV2"] diff --git a/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.py b/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.py new file mode 100644 index 00000000..6eb8cf9d --- /dev/null +++ b/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.py @@ -0,0 +1,110 @@ +"""Simplified personal memory summarizer using v2 memory tools.""" + +from ..base_memory_agent import BaseMemoryAgent +from ...core.context import C +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message, ToolCall +from ...core.utils import format_messages + + +@C.register_op() +class PersonalSummarizerV2(BaseMemoryAgent): + """Simplified personal memory summarizer that uses v2 memory tools. + + This summarizer follows a three-step workflow: + 1. AddMemoryDrafts: Generate initial memory drafts from context + 2. RetrieveRecentAndSimilarMemories: Retrieve similar and recent memories + 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( + **{ + "description": self.get_prompt("tool"), + "parameters": { + "type": "object", + "properties": { + "messages": { + "type": "array", + "items": { + "type": "object", + "properties": { + "role": { + "type": "string", + "description": "role", + }, + "content": { + "type": "string", + "description": "content", + }, + }, + "required": ["role", "content"], + }, + }, + }, + "required": ["messages"], + }, + }, + ) + + async def build_messages(self) -> list[Message]: + """Construct messages with context, memory_target, and memory_type information.""" + system_prompt = self.prompt_format( + prompt_name="system_prompt", + context=self.description + "\n" + format_messages(self.get_messages()), + memory_type=self.memory_type.value, + memory_target=self.memory_target, + ) + + messages = [ + Message(role=Role.SYSTEM, content=system_prompt), + Message(role=Role.USER, content=self.get_prompt("user_message")), + ] + return messages + + async def _reasoning_step(self, messages: list[Message], step: int, **kwargs) -> tuple[Message, bool]: + return await super()._reasoning_step(messages, step, **kwargs) + + async def _acting_step(self, assistant_message: Message, step: int, **kwargs) -> list[Message]: + """Execute tool calls with memory_target, memory_type, and author context.""" + messages: list[Message] = await super()._acting_step( + assistant_message, + step, + memory_type=self.memory_type.value, + memory_target=self.memory_target, + ref_memory_id=self.ref_memory_id, + author=self.author, + **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 + + 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 new file mode 100644 index 00000000..9a70d51b --- /dev/null +++ b/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2.yaml @@ -0,0 +1,63 @@ +tool: | + Extract and store personal memories from conversation context using a three-step workflow. + Use this tool to analyze dialogues and extract important personal information about users, + such as preferences, habits, personal background, relationships, and significant facts. + +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. + + **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: + {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. + + ## Your Tasks - Three-Step Workflow: + + ### 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. + + ### 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. + + ### 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`. + +user_message: | + Please analyze the context and update the memory store following the three-step workflow: + 1. First use `AddMemoryDrafts` to generate initial memory drafts + 2. Then use `RetrieveRecentAndSimilarMemories` to find related existing memories + 3. Finally use `UpdateMemories` to remove outdated memories and add new consolidated memories diff --git a/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2_simple.yaml b/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2_simple.yaml new file mode 100644 index 00000000..acc5bd35 --- /dev/null +++ b/reme_ai/mem_agent/summarizer_v2/personal_summarizer_v2_simple.yaml @@ -0,0 +1,53 @@ +tool: | + Extract and store personal memories from conversation context using a three-step workflow. + Use this tool to analyze dialogues and extract important personal information about users, + such as preferences, habits, personal background, relationships, and significant facts. + +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. + + **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: + {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. + + ## Your Tasks - Three-Step Workflow: + + ### Step 1: Generate Memory Drafts + Use the `AddMemoryDrafts` tool to produce a set of non-redundant, self-contained memory drafts that capture all important information explicitly stated in the context. Each draft should record ONE complete fact with accurate time metadata (year, month, day) when time references are mentioned. If no valuable information exists, output `` and stop. + + ### Step 2: Retrieve Similar and Recent Memories + Use the `RetrieveRecentAndSimilarMemories` tool to obtain all existing memories that are semantically related to each memory draft, ensuring comprehensive coverage for deduplication and conflict detection. + + ### Step 3: Update Memories + Use the `UpdateMemories` tool to produce a final, non-redundant memory set where: + - `memory_ids_to_delete` contains IDs of memories that are duplicates, outdated, or being consolidated + - `memories_to_add` contains new or updated memories that preserve all information without redundancy or conflicts + + ## Guidelines: + - **Be selective**: Store only truly important information. + - **Stay concise**: Each memory should be clear and atomic, recording ONE complete piece of information. + - **Be strictly accurate**: Ensure extracted content faithfully reflects ONLY what is explicitly stated in the original context. DO NOT infer, extrapolate, or fabricate any details. + - **AVOID REDUNDANCY AT ALL COSTS**: This is your TOP PRIORITY. Always perform thorough deduplication: + * First, deduplicate within newly extracted memory drafts + * Then, check against retrieved memories from Step 2 + * Actively use `memory_ids_to_delete` to remove duplicate or conflicting memories + * Use `memories_to_add` to consolidate and integrate information from multiple memories into one + * If information semantically matches existing memories, DO NOT add it again + * When uncertain, prefer to skip or update existing memories rather than create duplicates + - **Include relevant metadata**: Include time-related metadata (year, month, day) when appropriate, especially when time references are mentioned. + - **No assumptions**: Only store information that is directly and clearly stated in the conversation. + - **Quality over quantity**: It's better to have fewer, well-maintained memories than many duplicate ones. + +user_message: | + Please update the memory store following the three-step workflow. If there is no valuable information to remember, output `` without calling any tools. diff --git a/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.py b/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.py new file mode 100644 index 00000000..a580806a --- /dev/null +++ b/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.py @@ -0,0 +1,94 @@ +"""Simplified orchestrator for memory summarization workflow - V2.""" + +from loguru import logger + +from ..base_memory_agent import BaseMemoryAgent +from ...core.context import C +from ...core.enumeration import Role, MemoryType +from ...core.schema import Message, MemoryNode, ToolCall +from ...core.utils import format_messages + + +@C.register_op() +class ReMeSummarizerV2(BaseMemoryAgent): + """Simplified version that coordinates memory updates using only summary_and_hands_off tool.""" + + def __init__(self, meta_memories: list[dict] | None = None, **kwargs): + """Initialize with meta memories list.""" + super().__init__(**kwargs) + self.meta_memories: list[dict] = meta_memories or [] + + def _build_tool_call(self) -> ToolCall: + return ToolCall( + **{ + "description": self.prompt_format("tool"), + "parameters": { + "type": "object", + "properties": { + "messages": { + "type": "array", + "items": { + "type": "object", + "properties": { + "role": { + "type": "string", + "description": "role", + }, + "content": { + "type": "string", + "description": "content", + }, + }, + "required": ["role", "content"], + }, + }, + }, + "required": ["messages"], + }, + }, + ) + + async def _read_meta_memories(self) -> str: + """Fetch meta-memory entries using format_memory_metadata.""" + from ...mem_tool import ReadMetaMemory + + return ReadMetaMemory().format_memory_metadata(self.meta_memories) + + async def build_messages(self) -> list[Message]: + """Construct initial messages with context and meta-memory information.""" + messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages] + self.context["messages_formated"] = self.description + "\n" + format_messages(messages) + self.context["ref_memory_id"] = MemoryNode( + memory_type=MemoryType.HISTORY, + content=self.context["messages_formated"], + ).memory_id + + meta_memory_info = await self._read_meta_memories() + logger.info(f"meta_memory_info={meta_memory_info}") + + system_prompt = self.prompt_format( + prompt_name="system_prompt", + meta_memory_info=meta_memory_info, + context=self.context["messages_formated"], + ) + + user_message = self.get_prompt("user_message") + messages = [ + Message(role=Role.SYSTEM, content=system_prompt), + Message(role=Role.USER, content=user_message), + ] + + return messages + + async def _acting_step(self, assistant_message: Message, step: int, **kwargs) -> list[Message]: + """Execute tool calls with ref_memory_id and author context.""" + return await super()._acting_step( + assistant_message, + step, + messages=self.context.get("messages", []), + description=self.context.get("description"), + ref_memory_id=self.context["ref_memory_id"], + messages_formated=self.context["messages_formated"], + author=self.author, + **kwargs, + ) diff --git a/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.yaml b/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.yaml new file mode 100644 index 00000000..30792a08 --- /dev/null +++ b/reme_ai/mem_agent/summarizer_v2/reme_summarizer_v2.yaml @@ -0,0 +1,25 @@ +tool: | + Orchestrate the complete memory summarization for the agent. + +system_prompt: | + You are a Memory Agent responsible for performing necessary updates and summaries of the main Agent's memories based on the **context**. + + # Context + {context} + + ## Main Agent's Meta Memory + Each line of meta memory indicates the existence of a specialized Memory Agent dedicated to deep summarization and updating of memories within a specific dimension (memory_type + memory_target). + Format: "- (): " + {meta_memory_info} + + ## Your Task + Use `summary_and_hands_off` tool to: + 1. Create a concise summary in `summary_content` that captures key points, decisions, or important facts from the context. + 2. Identify which memory dimensions need updates and specify them in `memory_tasks` (each with `memory_type` and `memory_target`). + - The `memory_type` and `memory_target` must exactly match existing entries in the "Main Agent's Meta Memory" listed above. + - Multiple tasks can be specified to enable parallel processing by specialized agents. + + Note: If the context contains no memorable information (e.g., simple greetings), output ``. + +user_message: | + Please perform your task based on the context. diff --git a/reme_ai/mem_tool/base_memory_tool.py b/reme_ai/mem_tool/base_memory_tool.py index d0e6834e..2c2d4e50 100644 --- a/reme_ai/mem_tool/base_memory_tool.py +++ b/reme_ai/mem_tool/base_memory_tool.py @@ -33,9 +33,13 @@ class BaseMemoryTool(BaseOp, metaclass=ABCMeta): def _build_multiple_parameters(self) -> dict: return {} + def _build_tool_description(self) -> str: + """Build tool description.""" + return self.get_prompt("tool" + ("_multiple" if self.enable_multiple else "")) + def _build_tool_call(self) -> ToolCall: tool_call_params: dict = { - "description": self.get_prompt("tool" + ("_multiple" if self.enable_multiple else "")), + "description": self._build_tool_description(), } if self.enable_multiple: @@ -50,7 +54,7 @@ class BaseMemoryTool(BaseOp, metaclass=ABCMeta): parameters["properties"] = { "thinking": { "type": "string", - "description": "Your thinking and reasoning about how to fill in the parameters", + "description": "Your complete and detailed thinking process about how to fill in each parameter", }, **parameters["properties"], } @@ -78,6 +82,16 @@ class BaseMemoryTool(BaseOp, metaclass=ABCMeta): """Get the reference memory ID from context.""" return self.context.get("ref_memory_id", "") + @property + def messages_formated(self) -> str: + """Get the formated messages from context.""" + return self.context.get("messages_formated", "") + + @property + def retrieved_nodes(self) -> list[MemoryNode]: + """Get the retrieved nodes from context.""" + return self.context.get("retrieved_nodes") + @property def author(self) -> str: """Get the author from context.""" diff --git a/reme_ai/mem_tool/history/read_history_memory.py b/reme_ai/mem_tool/history/read_history_memory.py index d01cb4e8..def2ff24 100644 --- a/reme_ai/mem_tool/history/read_history_memory.py +++ b/reme_ai/mem_tool/history/read_history_memory.py @@ -43,7 +43,9 @@ class ReadHistoryMemory(BaseMemoryTool): ref_memory_id = self.context.get("ref_memory_id", "") ref_memory_ids: list[str] = [ref_memory_id] if ref_memory_id else [] + # Remove empty IDs and duplicates ref_memory_ids = [mid for mid in ref_memory_ids if mid] + ref_memory_ids = list(dict.fromkeys(ref_memory_ids)) # Remove duplicates while preserving order if not ref_memory_ids: self.output = "No valid reference memory IDs provided for reading." diff --git a/reme_ai/mem_tool/history/read_history_memory.yaml b/reme_ai/mem_tool/history/read_history_memory.yaml index 564af50f..24d99b18 100644 --- a/reme_ai/mem_tool/history/read_history_memory.yaml +++ b/reme_ai/mem_tool/history/read_history_memory.yaml @@ -8,4 +8,4 @@ ref_memory_id: | Reference memory ID to query the original history dialogue. ref_memory_ids: | - List of reference memory IDs to query the original history dialogues. + List of reference memory IDs to query the original history dialogues. Please provide unique IDs without duplicates. diff --git a/reme_ai/mem_tool/v2/__init__.py b/reme_ai/mem_tool/v2/__init__.py new file mode 100644 index 00000000..2ba9e900 --- /dev/null +++ b/reme_ai/mem_tool/v2/__init__.py @@ -0,0 +1,17 @@ +"""Version 2 memory tools with enhanced functionality.""" + +from .add_memory_drafts import AddMemoryDrafts +from .read_history import ReadHistory +from .retrieve_memories import RetrieveMemories +from .retrieve_recent_and_similar_memories import RetrieveRecentAndSimilarMemories +from .summary_and_hands_off import SummaryAndHandsOff +from .update_memories import UpdateMemories + +__all__ = [ + "AddMemoryDrafts", + "ReadHistory", + "RetrieveMemories", + "RetrieveRecentAndSimilarMemories", + "SummaryAndHandsOff", + "UpdateMemories", +] diff --git a/reme_ai/mem_tool/v2/add_memory_drafts.py b/reme_ai/mem_tool/v2/add_memory_drafts.py new file mode 100644 index 00000000..ba66c0e1 --- /dev/null +++ b/reme_ai/mem_tool/v2/add_memory_drafts.py @@ -0,0 +1,130 @@ +"""Add memory drafts operation for vector store.""" + +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ...core.context import C + + +@C.register_op() +class AddMemoryDrafts(BaseMemoryTool): + """Add memory drafts without persisting them to the database. + + This tool is useful for creating draft memories that can be reviewed and modified + before final submission. Drafts are not persisted to the vector store. + Metadata fields can be customized via `metadata_desc` parameter. + """ + + def __init__(self, add_when_to_use: bool = False, metadata_desc: dict[str, str] | None = None, **kwargs): + """Initialize AddMemoryDrafts. + + Args: + add_when_to_use: Include when_to_use field for better retrieval. Defaults to True. + metadata_desc: Dictionary defining metadata fields and their descriptions. + **kwargs: Additional arguments for BaseMemoryTool. + """ + kwargs["enable_multiple"] = True + super().__init__(**kwargs) + self.add_when_to_use: bool = add_when_to_use + self.metadata_desc: dict[str, str] = metadata_desc or {} + + def _build_item_schema(self) -> tuple[dict, list[str]]: + """Build shared schema properties and required fields for memory items to add. + + Returns: + Tuple of (properties dict, required fields list). + """ + properties = {} + required = [] + + if self.add_when_to_use: + properties["when_to_use"] = { + "type": "string", + "description": self.get_prompt("when_to_use"), + } + required.append("when_to_use") + + properties["memory_content"] = { + "type": "string", + "description": self.get_prompt("memory_content"), + } + required.append("memory_content") + + # Add metadata field if metadata_desc is provided and not empty + if self.metadata_desc: + metadata_properties = { + key: {"type": "string", "description": desc} for key, desc in self.metadata_desc.items() + } + properties["metadata"] = { + "type": "object", + "description": "metadata", + "properties": metadata_properties, + } + required.append("metadata") + + return properties, required + + def _build_multiple_parameters(self) -> dict: + """Build input schema for add drafts operation. + + Only supports batch mode for adding draft memories. + """ + item_properties, required_fields = self._build_item_schema() + return { + "type": "object", + "properties": { + "memory_drafts": { + "type": "array", + "description": self.get_prompt("memory_drafts"), + "items": { + "type": "object", + "properties": item_properties, + "required": required_fields, + }, + }, + }, + "required": ["memory_drafts"], + } + + def _extract_memory_data(self, mem_dict: dict) -> tuple[str, str, dict]: + """Extract memory data from a dictionary with proper defaults. + + Args: + mem_dict: Dictionary containing memory fields. + + Returns: + Tuple of (memory_content, when_to_use, metadata). + """ + memory_content = mem_dict.get("memory_content", "") + when_to_use = mem_dict.get("when_to_use", "") if self.add_when_to_use else "" + metadata = mem_dict.get("metadata", {}) if self.metadata_desc else {} + return memory_content, when_to_use, metadata + + 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) diff --git a/reme_ai/mem_tool/v2/add_memory_drafts.yaml b/reme_ai/mem_tool/v2/add_memory_drafts.yaml new file mode 100644 index 00000000..67809399 --- /dev/null +++ b/reme_ai/mem_tool/v2/add_memory_drafts.yaml @@ -0,0 +1,18 @@ +tool_multiple: | + Create draft memories for initial recording of information. + Use this tool to quickly capture information as drafts that can be reviewed or modified later. + **CRITICAL**: Only add memories based on explicitly stated facts. DO NOT store inferred, assumed, or fabricated information. + +memory_drafts: | + A list of draft memory objects to create. + Each draft represents a piece of information to be recorded initially. + +when_to_use: | + When to retrieve this memory. + This field is used for vector embedding to improve retrieval accuracy by providing contextual information. + +memory_content: | + The content of the memory draft to record. + Should be a clear, concise statement that captures the information to remember. + Keep it focused on a single piece of information for better retrieval accuracy. + **Must be strictly accurate and based only on explicitly stated facts - no inference or fabrication.** diff --git a/reme_ai/mem_tool/v2/read_history.py b/reme_ai/mem_tool/v2/read_history.py new file mode 100644 index 00000000..15aaada3 --- /dev/null +++ b/reme_ai/mem_tool/v2/read_history.py @@ -0,0 +1,57 @@ +"""Read history memory operation.""" + +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ...core.context import C +from ...core.schema import MemoryNode + + +@C.register_op() +class ReadHistory(BaseMemoryTool): + """Read original history dialogue by reference memory ID. + + Only supports single memory read (enable_multiple=False). + """ + + def __init__(self, **kwargs): + """Initialize ReadHistory. + + Args: + **kwargs: Additional args for BaseMemoryTool. + """ + # Force disable multiple mode + kwargs["enable_multiple"] = False + super().__init__(**kwargs) + + def _build_parameters(self) -> dict: + return { + "type": "object", + "properties": { + "ref_memory_id": { + "type": "string", + "description": self.get_prompt("ref_memory_id"), + }, + }, + "required": ["ref_memory_id"], + } + + async def execute(self): + ref_memory_id = self.context.get("ref_memory_id", "") + + if not ref_memory_id: + self.output = "No valid reference memory ID provided." + logger.warning(self.output) + return + + # Query history dialogue by ref_memory_id + nodes = await self.vector_store.get(vector_ids=[ref_memory_id]) + + if not nodes: + self.output = f"No history memory found with ID: {ref_memory_id}" + logger.warning(self.output) + return + + memory = MemoryNode.from_vector_node(nodes[0]) + self.output = memory.content + logger.info(f"Successfully read history memory: {ref_memory_id}") diff --git a/reme_ai/mem_tool/v2/read_history.yaml b/reme_ai/mem_tool/v2/read_history.yaml new file mode 100644 index 00000000..c79c79ad --- /dev/null +++ b/reme_ai/mem_tool/v2/read_history.yaml @@ -0,0 +1,5 @@ +tool: | + Read original history dialogue by reference memory ID. + +ref_memory_id: | + Reference memory ID to query the original history dialogue. diff --git a/reme_ai/mem_tool/v2/retrieve_memories.py b/reme_ai/mem_tool/v2/retrieve_memories.py new file mode 100644 index 00000000..abdfc377 --- /dev/null +++ b/reme_ai/mem_tool/v2/retrieve_memories.py @@ -0,0 +1,188 @@ +"""Retrieve memories using vector similarity search with multiple queries.""" + +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ...core.context import C +from ...core.schema import MemoryNode, VectorNode +from ...core.utils import deduplicate_memories + + +@C.register_op() +class RetrieveMemories(BaseMemoryTool): + """Retrieve memories using vector similarity search with multiple queries. + + Always requires memory_type/memory_target in the schema. + Only supports multiple query mode (enable_multiple=True). + Metadata filters can be customized via `metadata_desc` parameter for pre-retrieval filtering. + """ + + def __init__(self, metadata_desc: dict[str, str] | None = None, top_k: int = 20, **kwargs): + """Initialize RetrieveMemories. + + Args: + metadata_desc: Dictionary defining metadata filter fields and their descriptions. + These fields will be used as filters in vector search before similarity matching. + top_k: Max memories to retrieve per query. + **kwargs: Additional args for BaseMemoryTool. + """ + kwargs["enable_multiple"] = True + super().__init__(**kwargs) + self.metadata_desc: dict[str, str] = metadata_desc or {} + self.top_k: int = top_k + + def _build_query_schema(self) -> tuple[dict, list[str]]: + """Build schema properties and required fields for query items. + + Returns: + Tuple of (properties dict, required fields list). + """ + properties = { + "memory_type": { + "type": "string", + "description": self.get_prompt("memory_type"), + }, + "memory_target": { + "type": "string", + "description": self.get_prompt("memory_target"), + }, + "query": { + "type": "string", + "description": self.get_prompt("query"), + }, + } + required = ["memory_type", "memory_target", "query"] + + # Add metadata filter fields if metadata_desc is provided and not empty + if self.metadata_desc: + metadata_properties = { + key: {"type": "string", "description": desc} for key, desc in self.metadata_desc.items() + } + # Generate dynamic description based on metadata_desc fields + field_descriptions = "\n".join([f" - {key}: {desc}" for key, desc in self.metadata_desc.items()]) + metadata_description = ( + f"Optional metadata filters for narrowing search results. Available fields:\n{field_descriptions}" + ) + + properties["metadata_filters"] = { + "type": "object", + "description": metadata_description, + "properties": metadata_properties, + } + + return properties, required + + def _build_multiple_parameters(self) -> dict: + """Build input schema for multiple query mode. + + Returns: + Schema with query_items array. Each item has memory_type/memory_target/query. + """ + item_properties, item_required = self._build_query_schema() + return { + "type": "object", + "properties": { + "query_items": { + "type": "array", + "description": self.get_prompt("query_items"), + "items": { + "type": "object", + "properties": item_properties, + "required": item_required, + }, + }, + }, + "required": ["query_items"], + } + + async def _retrieve_by_query( + self, + memory_type: str, + memory_target: str, + query: str, + metadata_filters: dict | None = None, + ) -> list[MemoryNode]: + """Retrieve memories by query using vector similarity search. + + Args: + memory_type: Memory type to search. + memory_target: Memory target to search. + query: Query string for similarity search. + metadata_filters: Optional metadata filters to narrow search results. + + Returns: + List of matching memories. + """ + filter_dict = { + "memory_type": [memory_type], + "memory_target": [memory_target], + } + + # Add metadata filters if provided + if metadata_filters: + for key, value in metadata_filters.items(): + if value: # Only add non-empty filter values + value = str(value).strip() + filter_dict[key] = [value] if not isinstance(value, list) else value + + nodes: list[VectorNode] = await self.vector_store.search(query=query, limit=self.top_k, filters=filter_dict) + memory_nodes: list[MemoryNode] = [MemoryNode.from_vector_node(n) for n in nodes] + return memory_nodes + + async def execute(self): + """Execute memory retrieval based on multiple query items. + + Outputs formatted results or error message. + """ + query_items: list[dict] = self.context.get("query_items", []) + if not query_items: + self.output = "No query items provided for retrieval." + return + + # Filter out items without query text + query_items = [item for item in query_items if item.get("query")] + + if not query_items: + self.output = "No valid query texts provided for retrieval." + return + + # Retrieve memory_nodes for all queries + memory_nodes: list[MemoryNode] = [] + for item in query_items: + memory_type = item.get("memory_type") + memory_target = item.get("memory_target") + metadata_filters = item.get("metadata_filters", {}) if self.metadata_desc else {} + + if not memory_type or not memory_target: + logger.warning(f"Skipping query with missing memory_type or memory_target: {item}") + continue + + retrieved = await self._retrieve_by_query( + memory_type=memory_type, + memory_target=memory_target, + query=item["query"], + metadata_filters=metadata_filters, + ) + memory_nodes.extend(retrieved) + + # Deduplicate and format output + memory_nodes = deduplicate_memories(memory_nodes) + + # Build set of historical memory_ids for fast lookup + retrieved_memory_ids = {node.memory_id for node in self.retrieved_nodes if node.memory_id} + + # Filter out already retrieved memories by memory_id + new_memory_nodes = [node for node in memory_nodes if node.memory_id not in retrieved_memory_ids] + + # 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 memories found matching the queries (duplicates removed)." + else: + self.output = "\n".join([m.format_memory() for m in new_memory_nodes]) + + logger.info(f"Retrieved {len(memory_nodes)} memories, {len(new_memory_nodes)} new after deduplication") diff --git a/reme_ai/mem_tool/v2/retrieve_memories.yaml b/reme_ai/mem_tool/v2/retrieve_memories.yaml new file mode 100644 index 00000000..f83e11cb --- /dev/null +++ b/reme_ai/mem_tool/v2/retrieve_memories.yaml @@ -0,0 +1,24 @@ +tool_multiple: | + Retrieve memories from the memory store using multiple queries with vector similarity search. + Use this tool to find relevant memories based on semantic similarity to multiple queries. + This is useful when you need to search for different types of information in a single operation. + The search returns the most relevant memories ranked by similarity score for each query. + + Note: Within the same session, this tool automatically deduplicates results across multiple calls. + If you call this tool multiple times, only new memories (not previously retrieved) will be returned. + This prevents redundant information in subsequent retrievals. + +memory_type: | + The type of memory to search for. + You MUST select one of the memory_type values that are explicitly provided in the Available Meta-Memories. + +memory_target: | + The target of the memory to search within. + You MUST select one of the memory_type values that are explicitly provided in the Available Meta-Memories. + +query: | + The query text for vector similarity search. + Use descriptive queries that capture the semantic meaning of what you're looking for. + +query_items: | + A list of query items for vector similarity search. 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 new file mode 100644 index 00000000..5aba25ab --- /dev/null +++ b/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.py @@ -0,0 +1,176 @@ +"""Combined memory retrieval: recent + vector similarity search.""" + +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ...core.context import C +from ...core.schema import MemoryNode, VectorNode +from ...core.utils import deduplicate_memories + + +@C.register_op() +class RetrieveRecentAndSimilarMemories(BaseMemoryTool): + """Retrieve memories using both time-based and vector similarity search. + + First retrieves recent_top_k memories sorted by modification time, + then retrieves similar_top_k memories using vector similarity search. + Uses memory_type and memory_target from context (self.memory_type, self.memory_target). + """ + + def __init__( + self, + recent_top_k: int = 20, + similar_top_k: int = 20, + **kwargs, + ): + """Initialize RetrieveRecentAndSimilarMemories. + + Args: + recent_top_k: Max recent memories to retrieve by time. + similar_top_k: Max similar memories to retrieve by vector search. + **kwargs: Additional args for BaseMemoryTool. + """ + kwargs["enable_multiple"] = True + super().__init__(**kwargs) + self.recent_top_k: int = recent_top_k + self.similar_top_k: int = similar_top_k + + def _build_tool_description(self) -> str: + """Build tool description.""" + return self.prompt_format("tool_multiple", + recent_top_k=self.recent_top_k, + similar_top_k=self.similar_top_k) + + def _build_multiple_parameters(self) -> dict: + """Build input schema for multiple query mode. + + Returns: + Schema with query_items array. + """ + return { + "type": "object", + "properties": { + "query_items": { + "type": "array", + "description": self.get_prompt("query_items"), + "items": { + "type": "object", + "properties": { + "query": { + "type": "string", + "description": self.get_prompt("query"), + }, + }, + "required": ["query"], + }, + }, + }, + "required": ["query_items"], + } + + async def _retrieve_recent(self) -> list[MemoryNode]: + """Retrieve recent memories sorted by time_modified. + + Returns: + List of recent memories sorted by modification time (newest first). + """ + filter_dict = { + "memory_type": [self.memory_type.value], + "memory_target": [self.memory_target], + } + + # Use list() with sort_key="time_modified", reverse=True (descending), and limit + nodes: list[VectorNode] = await self.vector_store.list( + filters=filter_dict, + limit=self.recent_top_k, + sort_key="time_modified", + reverse=True, # Most recent first (descending order) + ) + + memory_nodes: list[MemoryNode] = [MemoryNode.from_vector_node(n) for n in nodes] + + return memory_nodes + + async def _retrieve_by_query( + self, + query: str, + ) -> list[MemoryNode]: + """Retrieve memories by query using vector similarity search. + + Args: + query: Query string for similarity search. + + Returns: + List of matching memories. + """ + filter_dict = { + "memory_type": [self.memory_type.value], + "memory_target": [self.memory_target], + } + + nodes: list[VectorNode] = await self.vector_store.search( + query=query, limit=self.similar_top_k, filters=filter_dict + ) + + memory_nodes: list[MemoryNode] = [MemoryNode.from_vector_node(n) for n in nodes] + + return memory_nodes + + async def execute(self): + """Execute combined memory retrieval (recent + similar). + + First retrieves recent_top_k memories by time, then retrieves similar_top_k + memories by vector similarity for each query in query_items. + Uses memory_type and memory_target from context. Outputs formatted results or error message. + """ + if not self.memory_type or not self.memory_target: + raise RuntimeError("memory_type and memory_target are required for retrieval.") + + # Get query items + query_items: list[dict] = self.context.get("query_items", []) + if not query_items: + self.output = "No query items provided for retrieval." + return + + # Filter out items without query text + query_items = [item for item in query_items if item.get("query")] + + if not query_items: + self.output = "No valid query texts provided for retrieval." + return + + # Step 1: Retrieve recent memories (once, shared across all queries) + recent_memory_nodes: list[MemoryNode] = await self._retrieve_recent() + logger.info(f"Retrieved {len(recent_memory_nodes)} recent memories") + + # Step 2: Retrieve similar memories by vector search for all queries + similar_memory_nodes: list[MemoryNode] = [] + for item in query_items: + retrieved = await self._retrieve_by_query(query=item["query"]) + similar_memory_nodes.extend(retrieved) + # Combine and deduplicate all memories + all_memory_nodes = recent_memory_nodes + similar_memory_nodes + all_memory_nodes = deduplicate_memories(all_memory_nodes) + + # Build set of historical memory_ids for fast lookup + retrieved_memory_ids = {node.memory_id for node in self.retrieved_nodes if node.memory_id} + + # Filter out already retrieved memories by memory_id + new_memory_nodes = [node for node in all_memory_nodes if node.memory_id not in retrieved_memory_ids] + + # 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: + self.output = "\n".join([m.format_memory() for m in new_memory_nodes]) + + logger.info( + f"Retrieved {len(all_memory_nodes)} total memories " + f"({len(recent_memory_nodes)} recent + {len(similar_memory_nodes)} similar), " + f"{len(new_memory_nodes)} new after deduplication" + ) diff --git a/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.yaml b/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.yaml new file mode 100644 index 00000000..ed91357e --- /dev/null +++ b/reme_ai/mem_tool/v2/retrieve_recent_and_similar_memories.yaml @@ -0,0 +1,25 @@ +tool_multiple: | + Retrieve memories using both time-based and multiple vector similarity searches. + + This tool combines two retrieval strategies: + 1. First retrieves the most recent memories based on modification time (recent top {recent_top_k}) + 2. Then retrieves semantically similar memories for each of your queries (similar top {similar_top_k} per query) + + This is useful when you need to search for different types of information in a single operation, + while also considering recent context. + + The results are automatically deduplicated, so you get a combined set of both recent + and relevant memories without duplicates. + + Note: Within the same session, this tool automatically deduplicates results across multiple calls. + If you call this tool multiple times, only new memories (not previously retrieved) will be returned. + This prevents redundant information in subsequent retrievals. + +query: | + The query text for vector similarity search. + Use descriptive queries that capture the semantic meaning of what you're looking for. + +query_items: | + A list of query items for vector similarity search. + Each query will be used to find semantically similar memories, which will be combined + with the recent memories retrieved based on modification time. diff --git a/reme_ai/mem_tool/v2/summary_and_hands_off.py b/reme_ai/mem_tool/v2/summary_and_hands_off.py new file mode 100644 index 00000000..4a595578 --- /dev/null +++ b/reme_ai/mem_tool/v2/summary_and_hands_off.py @@ -0,0 +1,158 @@ +"""Summary and hands-off tool for distributing summarized memory to appropriate agents.""" + +import json +from typing import TYPE_CHECKING + +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 + +if TYPE_CHECKING: + from ...mem_agent import BaseMemoryAgent + + +@C.register_op() +class SummaryAndHandsOff(BaseMemoryTool): + """Distribute summarized memory task to appropriate agent based on memory_type.""" + + def __init__(self, memory_agents: list["BaseMemoryAgent"], **kwargs): + # Force enable_multiple to True since this tool only supports multiple tasks + kwargs["enable_multiple"] = True + kwargs["sub_ops"] = memory_agents or [] + super().__init__(**kwargs) + from ...mem_agent import BaseMemoryAgent + + self.sub_ops: list[BaseMemoryAgent] = [a for a in self.sub_ops if isinstance(a, BaseMemoryAgent)] + + @property + def memory_agent_dict(self) -> dict[MemoryType, "BaseMemoryAgent"]: + """Returns a dictionary mapping memory types to their corresponding agents.""" + return {a.memory_type: a for a in self.sub_ops} + + def _build_item_schema(self) -> tuple[dict, list[str]]: + """Build shared schema properties and required fields for memory tasks.""" + properties = { + "memory_type": { + "type": "string", + "description": self.get_prompt("memory_type"), + "enum": [k.value for k in self.memory_agent_dict], + }, + "memory_target": { + "type": "string", + "description": self.get_prompt("memory_target"), + }, + } + required = ["memory_type", "memory_target"] + return properties, required + + def _build_multiple_parameters(self) -> dict: + """Build input schema for multiple summary and hands-off tasks.""" + item_properties, required_fields = self._build_item_schema() + return { + "type": "object", + "properties": { + "summary_content": { + "type": "string", + "description": self.get_prompt("summary_content"), + }, + "memory_tasks": { + "type": "array", + "description": self.get_prompt("memory_tasks"), + "items": { + "type": "object", + "properties": item_properties, + "required": required_fields, + }, + }, + }, + "required": ["summary_content", "memory_tasks"], + } + + @staticmethod + def _parse_memory_type_target(task: dict): + memory_type = task.get("memory_type", "") + memory_target = task.get("memory_target", "") + return {"memory_type": MemoryType(memory_type), "memory_target": memory_target} + + def _collect_tasks(self) -> list[dict]: + """Collect memory tasks from context.""" + tasks: list[dict] = [] + memory_tasks: list[dict] = self.context.get("memory_tasks", []) + for task in memory_tasks: + tasks.append(self._parse_memory_type_target(task)) + return tasks + + async def execute(self): + """Execute memory tasks by distributing to appropriate agents in parallel.""" + summary_content = self.context.get("summary_content", "") + assert summary_content, "No summary content provided." + + # Build and store summary node + summary_node = MemoryNode( + memory_type=MemoryType.HISTORY, + memory_target="", + when_to_use=summary_content, + content=self.messages_formated, + ref_memory_id="", + author=self.author, + metadata={}, + ) + logger.info(f"Adding summary node: {summary_node.model_dump_json(indent=2, exclude_none=True)}") + self.memory_nodes.append(summary_node) + vector_node = summary_node.to_vector_node() + await self.vector_store.delete(vector_ids=[vector_node.vector_id]) + await self.vector_store.insert([vector_node]) + + # Collect tasks + tasks = self._collect_tasks() + if not tasks: + self.output = "No valid memory tasks to execute." + return + + # Submit tasks to corresponding agents + agent_list = [] + for i, task in enumerate(tasks): + memory_type: MemoryType = task["memory_type"] + memory_target: str = task["memory_target"] + + if memory_type not in self.memory_agent_dict: + logger.warning(f"No agent found for memory_type={memory_type}") + continue + + agent = self.memory_agent_dict[memory_type].copy() + agent_list.append([agent, memory_type, memory_target]) + + logger.info(f"Task {i}: Submitting {memory_type.value} agent with summary for target={memory_target}") + self.submit_async_task( + agent.call, + query=self.context.get("query", ""), + messages=self.context.get("messages", []), + memory_type=memory_type, + memory_target=memory_target, + description=self.context.get("description"), + ref_memory_id=self.context.get("ref_memory_id", ""), + ) + + await self.join_async_tasks() + + # Collect results + results = [] + for i, (agent, memory_type, memory_target) in enumerate(agent_list): + result_str = str(agent.output) + if agent.memory_nodes: + self.memory_nodes.extend(agent.memory_nodes) + + results.append( + { + "memory_type": memory_type.value, + "memory_target": memory_target, + "result": result_str[:200] + ("..." if len(result_str) > 200 else ""), + } + ) + logger.info(f"Task {i}: Completed {memory_type.value} agent for target={memory_target}") + + results_str = json.dumps(results, ensure_ascii=False, indent=2) + self.output = f"Successfully executed summary and {len(results)} hands-off task(s):\n{results_str}" diff --git a/reme_ai/mem_tool/v2/summary_and_hands_off.yaml b/reme_ai/mem_tool/v2/summary_and_hands_off.yaml new file mode 100644 index 00000000..a7ca9598 --- /dev/null +++ b/reme_ai/mem_tool/v2/summary_and_hands_off.yaml @@ -0,0 +1,18 @@ +tool_multiple: | + Summarize and distribute memory tasks to appropriate agents in parallel. + Use this tool when you have already summarized the content and need to hand it off to specialized agents. + Each task will be processed by its corresponding memory agent based on memory_type. + +summary_content: | + The summarized content to be stored as memory. + Should be a clear, concise summary that captures the key information. + +memory_type: | + The type of memory to process. Determines which specialized agent handles the task. + +memory_target: | + The target entity for this memory. + This helps the agent focus on the specific subject of the memory task. + +memory_tasks: | + A list of memory tasks to distribute, each with memory_type and memory_target. diff --git a/reme_ai/mem_tool/v2/update_memories.py b/reme_ai/mem_tool/v2/update_memories.py new file mode 100644 index 00000000..48576e2e --- /dev/null +++ b/reme_ai/mem_tool/v2/update_memories.py @@ -0,0 +1,169 @@ +"""Update memories operation for vector store.""" + +from loguru import logger + +from ..base_memory_tool import BaseMemoryTool +from ...core.context import C +from ...core.schema import MemoryNode + + +@C.register_op() +class UpdateMemories(BaseMemoryTool): + """Update memories by removing old ones and adding new ones in a single atomic operation. + + This tool is useful for updating memories when you need to remove outdated information + and add updated information at the same time. Only supports batch mode (multiple operations). + Metadata fields can be customized via `metadata_desc` parameter. + """ + + def __init__(self, add_when_to_use: bool = False, metadata_desc: dict[str, str] | None = None, **kwargs): + """Initialize UpdateMemories. + + Args: + add_when_to_use: Include when_to_use field for better retrieval. Defaults to True. + metadata_desc: Dictionary defining metadata fields and their descriptions. + **kwargs: Additional arguments for BaseMemoryTool. + """ + kwargs["enable_multiple"] = True + super().__init__(**kwargs) + self.add_when_to_use: bool = add_when_to_use + self.metadata_desc: dict[str, str] = metadata_desc or {} + + def _build_item_schema(self) -> tuple[dict, list[str]]: + """Build shared schema properties and required fields for memory items to add. + + Returns: + Tuple of (properties dict, required fields list). + """ + properties = {} + required = [] + + if self.add_when_to_use: + properties["when_to_use"] = { + "type": "string", + "description": self.get_prompt("when_to_use"), + } + required.append("when_to_use") + + properties["memory_content"] = { + "type": "string", + "description": self.get_prompt("memory_content"), + } + required.append("memory_content") + + # Add metadata field if metadata_desc is provided and not empty + if self.metadata_desc: + metadata_properties = { + key: {"type": "string", "description": desc} for key, desc in self.metadata_desc.items() + } + properties["metadata"] = { + "type": "object", + "description": "metadata", + "properties": metadata_properties, + } + required.append("metadata") + + return properties, required + + def _build_multiple_parameters(self) -> dict: + """Build input schema for update operation. + + Only supports batch mode with both removal and addition. + """ + item_properties, required_fields = self._build_item_schema() + return { + "type": "object", + "properties": { + "memory_ids_to_delete": { + "type": "array", + "description": self.get_prompt("memory_ids_to_delete"), + "items": {"type": "string"}, + }, + "memories_to_add": { + "type": "array", + "description": self.get_prompt("memories_to_add"), + "items": { + "type": "object", + "properties": item_properties, + "required": required_fields, + }, + }, + }, + "required": ["memory_ids_to_delete", "memories_to_add"], + } + + def _extract_memory_data(self, mem_dict: dict) -> tuple[str, str, dict]: + """Extract memory data from a dictionary with proper defaults. + + Args: + mem_dict: Dictionary containing memory fields. + + Returns: + Tuple of (memory_content, when_to_use, metadata). + """ + memory_content = mem_dict.get("memory_content", "") + when_to_use = mem_dict.get("when_to_use", "") if self.add_when_to_use else "" + metadata = mem_dict.get("metadata", {}) if self.metadata_desc else {} + return memory_content, when_to_use, metadata + + async def execute(self): + """Execute update operation: first remove old memories by IDs, then add new updated memories.""" + # 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] + + # 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." + return + + removed_count = 0 + added_count = 0 + + # Step 1: Remove old memories + if memory_ids_to_delete: + await self.vector_store.delete(vector_ids=memory_ids_to_delete) + self.memory_nodes.extend(memory_ids_to_delete) + removed_count = len(memory_ids_to_delete) + logger.info(f"Removed {removed_count} memories from vector_store.") + + # Step 2: Add new updated memories + if memories_to_add: + memory_nodes: list[MemoryNode] = [] + for mem in memories_to_add: + memory_content, when_to_use, metadata = self._extract_memory_data(mem) + if not memory_content: + logger.warning("Skipping memory with empty content") + continue + + memory_nodes.append(self._build_memory_node(memory_content, when_to_use=when_to_use, metadata=metadata)) + + if memory_nodes: + # Convert to VectorNodes and collect IDs + vector_nodes = [node.to_vector_node() for node in memory_nodes] + vector_ids: list[str] = [node.vector_id for node in vector_nodes] + + # Delete existing IDs (upsert behavior), then insert + await self.vector_store.delete(vector_ids=vector_ids) + await self.vector_store.insert(nodes=vector_nodes) + added_count = len(memory_nodes) + logger.info(f"Added {added_count} new memories to vector_store.") + + self.memory_nodes.extend(memory_nodes) + + # Generate output message + operations = [] + if removed_count > 0: + operations.append(f"removed {removed_count} old memories") + if added_count > 0: + operations.append(f"added {added_count} new memories") + + if operations: + self.output = f"Successfully {' and '.join(operations)} in vector_store." + else: + self.output = "No valid operations performed. Please check your input." + + logger.info(self.output) diff --git a/reme_ai/mem_tool/v2/update_memories.yaml b/reme_ai/mem_tool/v2/update_memories.yaml new file mode 100644 index 00000000..717762af --- /dev/null +++ b/reme_ai/mem_tool/v2/update_memories.yaml @@ -0,0 +1,27 @@ +tool_multiple: | + Update memories by removing outdated ones and adding new ones in a single atomic operation. + Use this tool when you need to update the memory store by: + - Removing outdated or incorrect memories + - Adding new, updated information to replace the removed memories + - Performing a batch update where old memories are replaced with new, accurate information + Memory IDs for removal can be obtained from previous memory retrieval results. + **CRITICAL**: Only add memories based on explicitly stated facts. DO NOT store inferred, assumed, or fabricated information. + +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. + +memories_to_add: | + A list of new memory objects to add after removal. + These memories typically contain the updated information that replaces the removed memories. + +when_to_use: | + When to retrieve this memory. + This field is used for vector embedding to improve retrieval accuracy by providing contextual information. + +memory_content: | + The content of the memory to store. + Should be a clear, concise statement that captures the information to remember. + Keep it focused on a single piece of information for better retrieval accuracy. + **Must be strictly accurate and based only on explicitly stated facts - no inference or fabrication.** diff --git a/reme_ai/mem_tool/vector_store/add_memory.yaml b/reme_ai/mem_tool/vector_store/add_memory.yaml index 30df9476..9c712179 100644 --- a/reme_ai/mem_tool/vector_store/add_memory.yaml +++ b/reme_ai/mem_tool/vector_store/add_memory.yaml @@ -4,7 +4,8 @@ tool: | - Meta information: "I am very happy" - Personal preferences: "John prefers dark mode", "Alice works in PST timezone" - Procedural knowledge: "To deploy, run build then push", "Always validate input before processing" - - Tool usage tips: "search_tool works best with short queries", "Use cache tool for frequently accessed data" + + **CRITICAL**: Only add memories based on explicitly stated facts. DO NOT store inferred, assumed, or fabricated information. tool_multiple: | Add multiple memories to the vector store for future retrieval. @@ -12,19 +13,17 @@ tool_multiple: | Each memory can include when_to_use conditions and metadata for better organization and retrieval. Examples: storing multiple user preferences, multiple procedural steps, or multiple tool usage tips. + **CRITICAL**: Only add memories based on explicitly stated facts. DO NOT store inferred, assumed, or fabricated information in any of the memory entries. + when_to_use: | Optional condition description for when to retrieve this memory. This field is used for vector embedding to improve retrieval accuracy by providing contextual information. - Examples: - - "when user asks about authentication" - - "when deploying to production" - - "when using search_tool" - - "when handling error cases" memory_content: | The content of the memory to store. Should be a clear, concise statement that captures the information to remember. Keep it focused on a single piece of information for better retrieval accuracy. + **Must be strictly accurate and based only on explicitly stated facts - no inference or fabrication.** memories: | A list of memory objects to store. diff --git a/reme_ai/mem_tool/vector_store/add_summary_memory.py b/reme_ai/mem_tool/vector_store/add_summary_memory.py index a28a8956..ce1127ed 100644 --- a/reme_ai/mem_tool/vector_store/add_summary_memory.py +++ b/reme_ai/mem_tool/vector_store/add_summary_memory.py @@ -75,11 +75,11 @@ class AddSummaryMemory(AddMemory): ) -> MemoryNode: """Build MemoryNode from content, when_to_use, and metadata.""" node = MemoryNode( - memory_type=MemoryType.SUMMARY, + memory_type=MemoryType.HISTORY, memory_target="", - when_to_use="", - content=memory_content, - ref_memory_id=self.ref_memory_id, + when_to_use=memory_content, + content=self.messages_formated, + ref_memory_id="", author=self.author, metadata=metadata or {}, ) diff --git a/reme_ai/mem_tool/vector_store/add_summary_memory.yaml b/reme_ai/mem_tool/vector_store/add_summary_memory.yaml index 20c6a1f1..3b69bfc8 100644 --- a/reme_ai/mem_tool/vector_store/add_summary_memory.yaml +++ b/reme_ai/mem_tool/vector_store/add_summary_memory.yaml @@ -3,17 +3,7 @@ tool: | Use this tool to store a summarized version of the provided context. The LLM should first summarize the context, then call this tool with the summarized content. - This tool is specifically designed for storing summaries of conversations, events, or information - that has been condensed from a larger context. Examples: - - Summarizing a long conversation: "User discussed project requirements for a web app with authentication" - - Summarizing a decision: "Team decided to use PostgreSQL for the database after evaluating options" - - Summarizing an event: "Successfully deployed version 2.0 to production with new features" - summary_memory: | The summarized content to store as memory. Should be a clear, concise summary that captures the key information from the context. Keep it focused and informative - aim for 1-3 sentences that convey the essential points. - Examples: - - "User prefers Python for backend development and has experience with FastAPI framework" - - "Project deadline is January 15th, requires authentication, payment integration, and admin dashboard" - - "Bug in user registration was caused by missing email validation, fixed by adding regex check" diff --git a/reme_ai/mem_tool/vector_store/retrieve_recent_memory.py b/reme_ai/mem_tool/vector_store/retrieve_recent_memory.py index 43b683dd..0896d75a 100644 --- a/reme_ai/mem_tool/vector_store/retrieve_recent_memory.py +++ b/reme_ai/mem_tool/vector_store/retrieve_recent_memory.py @@ -16,17 +16,14 @@ class RetrieveRecentMemory(BaseMemoryTool): Uses memory_type and memory_target from context (self.memory_type, self.memory_target). """ - def __init__( - self, - top_k: int = 20, - **kwargs, - ): + def __init__(self, top_k: int = 20, **kwargs): """Initialize RetrieveRecentMemory. Args: top_k: Max memories to retrieve. **kwargs: Additional args for BaseMemoryTool. """ + kwargs["enable_multiple"] = False super().__init__(**kwargs) self.top_k: int = top_k @@ -60,19 +57,29 @@ class RetrieveRecentMemory(BaseMemoryTool): Outputs formatted results or error message. """ if not self.memory_type or not self.memory_target: - self.output = "memory_type and memory_target are required for retrieval." - return + raise RuntimeError("memory_type and memory_target are required for retrieval.") # Retrieve recent memories memory_nodes: list[MemoryNode] = await self._retrieve_recent() # Deduplicate and format output memory_nodes = deduplicate_memories(memory_nodes) - self.memory_nodes = memory_nodes - if not memory_nodes: - self.output = "No memory_nodes found." + # Build set of historical memory_ids for fast lookup + retrieved_memory_ids = {node.memory_id for node in self.retrieved_nodes if node.memory_id} + + # Filter out already retrieved memories by memory_id + new_memory_nodes = [node for node in memory_nodes if node.memory_id not in retrieved_memory_ids] + + # 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: - self.output = "\n".join([m.format_memory() for m in memory_nodes]) + self.output = "\n".join([m.format_memory() for m in new_memory_nodes]) - logger.info(f"Retrieved {len(memory_nodes)} recent memory_nodes") + logger.info(f"Retrieved {len(memory_nodes)} memory_nodes, {len(new_memory_nodes)} new after deduplication") diff --git a/reme_ai/mem_tool/vector_store/update_memory.yaml b/reme_ai/mem_tool/vector_store/update_memory.yaml index 9f01c42e..e156f2e7 100644 --- a/reme_ai/mem_tool/vector_store/update_memory.yaml +++ b/reme_ai/mem_tool/vector_store/update_memory.yaml @@ -7,6 +7,8 @@ tool: | - You need to refine or improve the clarity of stored information Memory ID can be obtained from previous memory retrieval results. + **CRITICAL**: Only update memories with explicitly stated facts. DO NOT introduce inferred, assumed, or fabricated information in the updated content. + tool_multiple: | Update multiple memories in the vector store by replacing old memories with new content. Use this tool for batch updates when: @@ -16,6 +18,8 @@ tool_multiple: | - You need to refine or improve multiple stored memories at once Memory IDs can be obtained from previous memory retrieval results. + **CRITICAL**: Only update memories with explicitly stated facts. DO NOT introduce inferred, assumed, or fabricated information in any of the updated contents. + memory_id: | The unique identifier (memory_id) of the old memory to be replaced. This ID is returned when memories are retrieved or added. @@ -28,6 +32,7 @@ memory_content: | The new content of the memory to store. Should be a clear, concise statement that captures the updated information to remember. Keep it focused on a single piece of information for better retrieval accuracy. + **Must be strictly accurate and based only on explicitly stated facts - no inference or fabrication.** memories: | A list of memory update objects. diff --git a/reme_ai/mem_tool/vector_store/vector_retrieve_memory.py b/reme_ai/mem_tool/vector_store/vector_retrieve_memory.py index 2d758994..655de36e 100644 --- a/reme_ai/mem_tool/vector_store/vector_retrieve_memory.py +++ b/reme_ai/mem_tool/vector_store/vector_retrieve_memory.py @@ -237,11 +237,22 @@ class VectorRetrieveMemory(BaseMemoryTool): # Deduplicate and format output memory_nodes = deduplicate_memories(memory_nodes) - self.memory_nodes = memory_nodes - if not memory_nodes: - self.output = "No memory_nodes found matching the query." + # Build set of historical memory_ids for fast lookup + retrieved_memory_ids = {node.memory_id for node in self.retrieved_nodes if node.memory_id} + + # Filter out already retrieved memories by memory_id + new_memory_nodes = [node for node in memory_nodes if node.memory_id not in retrieved_memory_ids] + + # 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 matching the query (duplicates removed)." else: - self.output = "\n".join([m.format_memory() for m in memory_nodes]) + self.output = "\n".join([m.format_memory() for m in new_memory_nodes]) - logger.info(f"Retrieved {len(memory_nodes)} memory_nodes") + logger.info(f"Retrieved {len(memory_nodes)} memory_nodes, {len(new_memory_nodes)} new after deduplication") diff --git a/reme_ai/mem_tool/vector_store/vector_retrieve_memory.yaml b/reme_ai/mem_tool/vector_store/vector_retrieve_memory.yaml index 4ab56f9f..f0752f4a 100644 --- a/reme_ai/mem_tool/vector_store/vector_retrieve_memory.yaml +++ b/reme_ai/mem_tool/vector_store/vector_retrieve_memory.yaml @@ -3,25 +3,29 @@ tool: | Use this tool to find relevant memories based on semantic similarity to the query. The search returns the most relevant memories ranked by similarity score. + Note: Within the same session, this tool automatically deduplicates results across multiple calls. + If you call this tool multiple times, only new memories (not previously retrieved) will be returned. + This prevents redundant information in subsequent retrievals. + tool_multiple: | Retrieve memories from the memory store using multiple queries with vector similarity search. Use this tool to find relevant memories based on semantic similarity to multiple queries. This is useful when you need to search for different types of information in a single operation. The search returns the most relevant memories ranked by similarity score for each query. + Note: Within the same session, this tool automatically deduplicates results across multiple calls. + If you call this tool multiple times, only new memories (not previously retrieved) will be returned. + This prevents redundant information in subsequent retrievals. + memory_type: | The type of memory to search for. Must be one of: - "identity": Information about the AI agent's identity, role, or characteristics - "personal": Information about users, their preferences, or personal details - - "procedural": Step-by-step instructions, workflows, or how-to knowledge - - "tool": Tool usage tips, examples, and best practices memory_target: | The target of the memory to search within. - For "personal" memory: the person's name or identifier (e.g., "john", "alice") - For "procedural" memory: the process or task name (e.g., "deployment", "authentication") - - For "tool" memory: the tool name (e.g., "search_tool", "calculator") - - For "identity" memory: typically "self" or the agent's identifier query: | The query text for vector similarity search. diff --git a/reme_ai/reme.py b/reme_ai/reme.py index a4f9de4c..3aed1afa 100644 --- a/reme_ai/reme.py +++ b/reme_ai/reme.py @@ -7,21 +7,32 @@ from .core.embedding import BaseEmbeddingModel from .core.enumeration import Role from .core.llm import BaseLLM from .core.schema import Message +from .core.utils import singleton from .core.vector_store import BaseVectorStore from .mem_agent.retriever import ReMeRetriever +from .mem_agent.retriever_v2 import ReMeRetrieverV2 from .mem_agent.summarizer import ReMeSummarizer, PersonalSummarizer +from .mem_agent.summarizer_v2 import ReMeSummarizerV2, PersonalSummarizerV2 from .mem_tool import ( HandsOffTool, ReadHistoryMemory, - AddMetaMemory, AddMemory, AddSummaryMemory, DeleteMemory, UpdateMemory, VectorRetrieveMemory, ) +from .mem_tool.v2 import ( + AddMemoryDrafts, + ReadHistory, + RetrieveMemories, + RetrieveRecentAndSimilarMemories, + SummaryAndHandsOff, + UpdateMemories, +) +@singleton class ReMe(Application): """Simplified ReMe application that auto-initializes the service context.""" @@ -62,6 +73,10 @@ class ReMe(Application): self.vector_store: BaseVectorStore = C.get_vector_store("default") self.embedding_model: BaseEmbeddingModel = C.get_embedding_model("default") + @staticmethod + def get_llm(name: str) -> BaseLLM: + return C.get_llm(name) + @staticmethod def _prepare_messages(messages: list[dict | Message], user_id: str, assistant_id: str): if not messages: @@ -92,7 +107,7 @@ class ReMe(Application): "year": "The `year` information associated with the memory(Optional)", "month": "The `month` information associated with the memory(Optional)", "day": "The `day` information associated with the memory(Optional)", - "hour": "The `hour` information associated with the memory(Optional)", + # "hour": "The `hour` information associated with the memory(Optional)", # "year": "The year when the memory content occurred(Optional)", # "month": "The month when the memory content occurred(Optional)", # "day": "The day when the memory content occurred(Optional)", @@ -110,7 +125,6 @@ class ReMe(Application): personal_summarizer = PersonalSummarizer( tools=[ VectorRetrieveMemory( - enable_summary_memory=False, add_memory_type_target=False, metadata_desc=None, top_k=15, @@ -125,14 +139,19 @@ class ReMe(Application): meta_memories=meta_memories, enable_identity_memory=False, tools=[ - AddMetaMemory(), + # AddMetaMemory(), AddSummaryMemory(metadata_desc=metadata_summary), HandsOffTool(memory_agents=[personal_summarizer]), ], ) - await reme_summarizer.call(messages=messages, description=description, **kwargs) - return reme_summarizer.memory_nodes + try: + await reme_summarizer.call(messages=messages, description=description, **kwargs) + return reme_summarizer.memory_nodes + except Exception as e: + print(f"Warning: reme_summarizer.call failed: {e}") + return [] + else: raise NotImplementedError @@ -144,6 +163,7 @@ class ReMe(Application): description: str = "", user_id: str = "", assistant_id: str = "", + top_k: int = 20, **kwargs, ): """Retrieves relevant memories based on the query and specified memory mode.""" @@ -155,7 +175,7 @@ class ReMe(Application): "year": "The year to filter memories(Optional)", "month": "The month to filter memories(Optional)", "day": "The day to filter memories(Optional)", - "hour": "The hour to filter memories(Optional)", + # "hour": "The hour to filter memories(Optional)", } meta_memories = [ { @@ -168,18 +188,129 @@ class ReMe(Application): meta_memories=meta_memories, tools=[ VectorRetrieveMemory( - enable_summary_memory=True, add_memory_type_target=True, metadata_desc=metadata_retrieve, - top_k=20, + top_k=top_k, ), ReadHistoryMemory(), ], ) - await reme_retriever.call(query=query, messages=messages, description=description, **kwargs) - - return reme_retriever.output + try: + await reme_retriever.call(query=query, messages=messages, description=description, **kwargs) + return reme_retriever.output + except Exception as e: + print(f"Warning: reme_retriever.call failed: {e}") + return "error, not retrieved" + + + else: + raise NotImplementedError + + async def summary_v2( + self, + messages: list[dict], + description: str = "", + user_id: str = "", + assistant_id: str = "", + **kwargs, + ): + """Summarizes messages using V2 workflow with simplified tools.""" + + 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.", + } + meta_memories = [ + { + "memory_type": "personal", + "memory_target": user_id, + }, + ] + + messages = self._prepare_messages(messages, user_id, assistant_id) + + personal_summarizer_v2 = PersonalSummarizerV2( + tools=[ + AddMemoryDrafts(enable_thinking_params=True, metadata_desc=metadata_desc), + RetrieveRecentAndSimilarMemories( + enable_thinking_params=True, + metadata_desc=None, + recent_top_k=20, + similar_top_k=20, + ), + UpdateMemories(enable_thinking_params=True, metadata_desc=metadata_desc), + ], + ) + + reme_summarizer_v2 = ReMeSummarizerV2( + meta_memories=meta_memories, + tools=[ + SummaryAndHandsOff( + enable_thinking_params=True, + metadata_desc=metadata_desc, + memory_agents=[personal_summarizer_v2], + ), + ], + ) + + 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 + + else: + raise NotImplementedError + + async def retrieve_v2( + self, + query: str = "", + messages: list[dict] | None = None, + description: str = "", + user_id: str = "", + assistant_id: str = "", + top_k: int = 20, + **kwargs, + ): + """Retrieves relevant memories using V2 workflow with autonomous retrieval.""" + + if user_id: + messages = self._prepare_messages(messages, user_id, assistant_id) + + metadata_retrieve = { + "year": "The year to filter memories.", + "month": "The month to filter memories.", + "day": "The day to filter memories.", + } + meta_memories = [ + { + "memory_type": "personal", + "memory_target": user_id, + }, + ] + + reme_retriever_v2 = ReMeRetrieverV2( + meta_memories=meta_memories, + tools=[ + RetrieveMemories( + enable_thinking_params=True, + metadata_desc=metadata_retrieve, + top_k=top_k, + ), + # ReadHistory(enable_thinking_params=True), + ], + ) + + 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 diff --git a/tests/test_reme.py b/tests/test_reme.py index 2e57e8a7..c9e6843e 100644 --- a/tests/test_reme.py +++ b/tests/test_reme.py @@ -13,7 +13,7 @@ reme = ReMe( async def test_reme(): """Tests ReMe memory system with personal information storage and retrieval.""" # 构建一段包含个人信息的对话 - await reme.vector_store.delete_collection("reme") + await reme.vector_store.delete_all() messages = [ { @@ -55,7 +55,8 @@ async def test_reme(): print("=" * 60) # 对对话进行总结,生成记忆 - await reme.summary( + # await reme.summary( + await reme.summary_v2( messages=messages, user_id="zhangwei", description="用户自我介绍和技术兴趣分享", @@ -81,25 +82,25 @@ async def test_reme(): # 测试问题1: 检索用户姓名 query1 = "用户叫什么名字?" print(f"\n问题1: {query1}") - result1 = await reme.retrieve(query=query1, user_id="zhangwei") + result1 = await reme.retrieve_v2(query=query1, user_id="zhangwei") print(f"检索结果:\n{result1}") # 测试问题2: 检索技术背景 query2 = "用户擅长什么编程语言和技术方向?" print(f"\n问题2: {query2}") - result2 = await reme.retrieve(query=query2, user_id="zhangwei") + result2 = await reme.retrieve_v2(query=query2, user_id="zhangwei") print(f"检索结果:\n{result2}") # 测试问题3: 检索个人信息 query3 = "用户的工作地点和联系方式是什么?" print(f"\n问题3: {query3}") - result3 = await reme.retrieve(query=query3, user_id="zhangwei") + result3 = await reme.retrieve_v2(query=query3, user_id="zhangwei") print(f"检索结果:\n{result3}") # 测试问题4: 检索兴趣爱好 query4 = "用户平时有什么爱好或活动?" print(f"\n问题4: {query4}") - result4 = await reme.retrieve(query=query4, user_id="zhangwei") + result4 = await reme.retrieve_v2(query=query4, user_id="zhangwei") print(f"检索结果:\n{result4}") print("\n" + "=" * 60)