feat(benchmark): add HaluMem baseline evaluation and analysis tools

This commit is contained in:
jinli.yl 2026-01-13 23:59:55 +08:00
parent 6125dff01e
commit b8124fe31a
14 changed files with 1430 additions and 540 deletions

View file

@ -0,0 +1,180 @@
"""
分析 bench_results/reme_simple/tmp 目录下的评估结果
统计所有用户session中的result_type分布,并输出非Correct结果的详细位置信息。
"""
import json
from collections import Counter
from pathlib import Path
from typing import Dict, List, Tuple
def analyze_results(tmp_dir: str = "bench_results/reme_simple/tmp"):
"""
分析评估结果目录。
Args:
tmp_dir: 临时结果目录路径
"""
tmp_path = Path(tmp_dir)
if not tmp_path.exists():
print(f"❌ 目录不存在: {tmp_dir}")
return
# 统计数据
result_counter = Counter()
non_correct_results = [] # 存储非Correct结果的详细信息
# 遍历所有用户目录
user_dirs = sorted([d for d in tmp_path.iterdir() if d.is_dir()])
if not user_dirs:
print(f"❌ {tmp_dir} 下没有用户目录")
return
print(f"📁 找到 {len(user_dirs)} 个用户目录\n")
print("=" * 80)
print("开始分析...")
print("=" * 80 + "\n")
total_sessions = 0
total_questions = 0
# 遍历每个用户目录
for user_dir in user_dirs:
user_name = user_dir.name
# 获取该用户的所有session文件
session_files = sorted([
f for f in user_dir.iterdir()
if f.name.startswith("session_") and f.suffix == ".json"
])
if not session_files:
continue
# 遍历每个session
for session_file in session_files:
try:
with open(session_file, "r", encoding="utf-8") as f:
session_data = json.load(f)
session_id = session_data.get("session_id", -1)
total_sessions += 1
# 跳过生成的QA session
if session_data.get("is_generated_qa_session", False):
continue
# 获取评估结果
eval_results = session_data.get("evaluation_results", {})
qa_records = eval_results.get("question_answering_records", [])
# 分析每个问题的结果
for qa_idx, qa_record in enumerate(qa_records):
result_type = qa_record.get("result_type", "Unknown")
# 统计result_type
result_counter[result_type] += 1
total_questions += 1
# 如果不是Correct,记录详细信息
if result_type != "Correct":
non_correct_results.append({
"user_name": user_name,
"session_id": session_id,
"question_id": qa_idx,
"result_type": result_type,
"question": qa_record.get("question", ""),
"answer": qa_record.get("answer", ""),
"system_response": qa_record.get("system_response", "")
})
except Exception as e:
print(f"⚠️ 读取文件失败: {session_file}, 错误: {e}")
continue
# 输出统计结果
print("\n" + "=" * 80)
print("统计结果")
print("=" * 80 + "\n")
print(f"📊 总用户数: {len(user_dirs)}")
print(f"📊 总Session数: {total_sessions}")
print(f"📊 总问题数: {total_questions}\n")
if total_questions == 0:
print("❌ 没有找到任何问题数据")
return
# 输出result_type分布
print("=" * 80)
print("Result Type 分布")
print("=" * 80 + "\n")
# 按数量降序排列
sorted_results = sorted(result_counter.items(), key=lambda x: x[1], reverse=True)
for result_type, count in sorted_results:
ratio = count / total_questions * 100
print(f" {result_type:20s}: {count:5d} ({ratio:6.2f}%)")
# 输出非Correct结果的详细信息
if non_correct_results:
print("\n" + "=" * 80)
print(f"非 Correct 结果详情 (共 {len(non_correct_results)} 条)")
print("=" * 80 + "\n")
for idx, result in enumerate(non_correct_results, 1):
print(f"[{idx}] {result['result_type']}")
print(f" 用户: {result['user_name']}")
print(f" 位置: Session {result['session_id']}, Question {result['question_id']}")
print(f" 问题: {result['question']}")
print(f" 正确答案: {result['answer']}")
print(f" 系统回答: {result['system_response'][:200]}{'...' if len(result['system_response']) > 200 else ''}")
print()
else:
print("\n🎉 所有问题都是 Correct!")
# 保存详细报告到文件
report_file = Path(tmp_dir).parent / "analysis_report.json"
report_data = {
"summary": {
"total_users": len(user_dirs),
"total_sessions": total_sessions,
"total_questions": total_questions,
"result_type_distribution": dict(result_counter),
"result_type_ratio": {
result_type: count / total_questions
for result_type, count in result_counter.items()
}
},
"non_correct_results": non_correct_results
}
with open(report_file, "w", encoding="utf-8") as f:
json.dump(report_data, f, ensure_ascii=False, indent=2)
print("=" * 80)
print(f"📄 详细报告已保存到: {report_file}")
print("=" * 80)
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(
description="分析 ReMe 评估结果中的 result_type 分布"
)
parser.add_argument(
"--tmp_dir",
type=str,
default="bench_results/reme_simple/tmp",
help="临时结果目录路径 (默认: bench_results/reme_simple/tmp)"
)
args = parser.parse_args()
analyze_results(args.tmp_dir)

View file

@ -0,0 +1,604 @@
"""
HaluMem Benchmark Evaluator - Baseline (Direct QA without Memory System)
A simple baseline evaluation pipeline that:
1. Loads HaluMem benchmark data
2. Directly uses dialogue history to answer questions (no memory system)
3. Evaluates question answering performance
4. Generates comprehensive metrics
Usage:
python bench/halumem/eval_baseline_simple.py \
--data_path /path/to/HaluMem-Medium.jsonl \
--user_num 100 --max_concurrency 20
"""
import asyncio
import json
import os
import re
import time
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
from loguru import logger
from eval_tools import evaluation_for_question2
from llms import llm_request_for_json
# ==================== Configuration ====================
@dataclass
class EvalConfig:
"""Evaluation configuration parameters."""
data_path: str
user_num: int = 1
max_concurrency: int = 2
output_dir: str = "bench_results/baseline_simple"
# ==================== Utilities ====================
class DataLoader:
"""Handles loading and parsing of HaluMem data."""
@staticmethod
def load_jsonl(file_path: str) -> list[dict]:
"""Load all entries from a JSONL file."""
with open(file_path, "r", encoding="utf-8") as f:
return [json.loads(line.strip()) for line in f if line.strip()]
@staticmethod
def extract_user_name(persona_info: str) -> str:
"""Extract user name from persona info string."""
match = re.search(r"Name:\s*(.*?); Gender:", persona_info)
if not match:
raise ValueError(f"No name found in persona_info: {persona_info}")
return match.group(1).strip()
@staticmethod
def format_dialogue_messages(dialogue: list[dict]) -> list[dict]:
"""Format dialogue into ReMe message format."""
return [
{
"role": turn["role"],
"content": turn["content"],
"time_created": datetime.strptime(
turn["timestamp"], "%b %d, %Y, %H:%M:%S"
)
.replace(tzinfo=timezone.utc)
.strftime("%Y-%m-%d %H:%M:%S"),
}
for turn in dialogue
]
@staticmethod
def format_dialogue_for_eval(dialogue: list[dict], user_name: str = None) -> str:
"""Format dialogue into string for evaluation."""
formatted_turns = []
for turn in dialogue:
timestamp = datetime.strptime(
turn["timestamp"], "%b %d, %Y, %H:%M:%S"
).replace(tzinfo=timezone.utc).strftime("%Y-%m-%d %H:%M:%S")
# Use user_name if role is 'user' and user_name is provided
role = user_name if turn['role'] == 'user' and user_name else turn['role']
formatted_turns.append(
f"Role: {role}\n"
f"Content: {turn['content']}\n"
f"Time: {timestamp}"
)
return "\n\n".join(formatted_turns)
class FileManager:
"""Manages file I/O operations."""
def __init__(self, base_dir: str):
self.base_dir = Path(base_dir)
self.tmp_dir = self.base_dir / "tmp"
self.tmp_dir.mkdir(parents=True, exist_ok=True)
def get_user_dir(self, user_name: str) -> Path:
"""Get the directory path for a user."""
user_dir = self.tmp_dir / user_name
user_dir.mkdir(parents=True, exist_ok=True)
return user_dir
def get_session_file(self, user_name: str, session_id: int) -> Path:
"""Get the file path for a specific session."""
return self.get_user_dir(user_name) / f"session_{session_id}.json"
def save_session(self, user_name: str, session_id: int, data: dict):
"""Save session data to file."""
file_path = self.get_session_file(user_name, session_id)
with open(file_path, "w", encoding="utf-8") as f:
json.dump(data, f, ensure_ascii=False, indent=2)
logger.info(f"✅ Saved session {session_id} to {file_path}")
def load_session(self, user_name: str, session_id: int) -> dict | None:
"""Load session data from file."""
file_path = self.get_session_file(user_name, session_id)
if not file_path.exists():
return None
with open(file_path, "r", encoding="utf-8") as f:
return json.load(f)
def user_has_cache(self, user_name: str) -> bool:
"""Check if user has cached results."""
user_dir = self.get_user_dir(user_name)
return any(f.name.startswith("session_") and f.suffix == ".json"
for f in user_dir.iterdir())
def combine_results(self, output_file: str):
"""Combine all user session files into a single JSONL file."""
with open(output_file, "w", encoding="utf-8") as f_out:
for user_dir in self.tmp_dir.iterdir():
if not user_dir.is_dir():
continue
session_files = sorted([
f for f in user_dir.iterdir()
if f.name.startswith("session_") and f.suffix == ".json"
])
if not session_files:
continue
# Load first session to get user metadata
with open(session_files[0], "r", encoding="utf-8") as f_in:
first_session = json.load(f_in)
user_data = {
"uuid": first_session["uuid"],
"user_name": first_session["user_name"],
"sessions": []
}
# Load all sessions
for session_file in session_files:
with open(session_file, "r", encoding="utf-8") as f_in:
session_data = json.load(f_in)
# Remove redundant user metadata
session_data.pop("uuid", None)
session_data.pop("user_name", None)
user_data["sessions"].append(session_data)
f_out.write(json.dumps(user_data, ensure_ascii=False) + "\n")
# ==================== Question Answering Prompt ====================
BASELINE_QA_PROMPT = """You are a helpful AI assistant. Based on the dialogue history provided below, please answer the question.
**Dialogue History:**
{dialogue}
**Question:**
{question}
**Instructions:**
- Carefully read through the dialogue history
- Answer the question based ONLY on information present in the dialogue
- If the information needed to answer the question is NOT in the dialogue, respond with "I don't know" or "The information is not available in the dialogue"
- Do NOT make up or hallucinate information that is not explicitly mentioned in the dialogue
- Provide your reasoning process before giving the final answer
**Response Format:**
Please respond in JSON format with the following structure:
```json
{{
"reasoning": "Your step-by-step reasoning process",
"answer": "Your final answer (or 'I don't know' if information is not available)"
}}
```"""
# ==================== Evaluation ====================
class BaselineQuestionAnsweringEvaluator:
"""Evaluates question answering performance using direct LLM inference (no memory system)."""
def __init__(self):
pass
async def answer_question(
self,
question: str,
formatted_dialogue: str
) -> tuple[str, str, float]:
"""
Answer a question using the dialogue history directly.
Returns:
tuple: (answer, reasoning, duration_ms)
"""
start = time.time()
# Format prompt
prompt = BASELINE_QA_PROMPT.format(
dialogue=formatted_dialogue,
question=question
)
# Get answer from LLM
try:
# model_name = "qwen3-max"
model_name = "qwen3-30b-a3b-instruct-2507"
result = await llm_request_for_json(prompt, model_name=model_name)
answer = result.get("answer", "I don't know")
reasoning = result.get("reasoning", "")
except Exception as e:
logger.error(f"Error getting answer from LLM: {e}")
answer = "Error: Failed to get answer"
reasoning = str(e)
duration_ms = (time.time() - start) * 1000
return answer, reasoning, duration_ms
async def evaluate_questions(
self,
questions: list[dict],
user_name: str,
uuid: str,
session_id: int,
formatted_dialogue: str
) -> list[dict]:
"""Evaluate all questions for a session."""
results = []
for qa in questions:
# Get answer directly from LLM
answer, reasoning, duration_ms = await self.answer_question(
question=qa["question"],
formatted_dialogue=formatted_dialogue
)
# Evaluate response
evidence_text = "\n".join([e["memory_content"] for e in qa["evidence"]])
eval_result = await evaluation_for_question2(
qa["question"],
qa["answer"],
evidence_text,
answer,
formatted_dialogue
)
# Build result record
qa_result = {
**qa,
"uuid": uuid,
"session_id": session_id,
"system_response": answer,
"reasoning": reasoning,
"answer_duration_ms": duration_ms,
"result_type": eval_result.get("evaluation_result"),
"question_answering_reasoning": eval_result.get("reasoning", "")
}
results.append(qa_result)
return results
class MetricsAggregator:
"""Aggregates evaluation metrics."""
@staticmethod
def compute_qa_metrics(qa_records: list[dict]) -> dict[str, Any]:
"""Compute question answering metrics."""
total = len(qa_records)
if total == 0:
return {
"correct_qa_ratio(all)": 0,
"hallucination_qa_ratio(all)": 0,
"omission_qa_ratio(all)": 0,
"correct_qa_ratio(valid)": 0,
"hallucination_qa_ratio(valid)": 0,
"omission_qa_ratio(valid)": 0,
"qa_valid_num": 0,
"qa_num": 0
}
correct = 0
hallucination = 0
omission = 0
valid = 0
for qa in qa_records:
result_type = qa.get("result_type", "")
if result_type in ["Correct", "Hallucination", "Omission"]:
valid += 1
if result_type == "Correct":
correct += 1
elif result_type == "Hallucination":
hallucination += 1
elif result_type == "Omission":
omission += 1
metrics = {
"correct_qa_ratio(all)": correct / total,
"hallucination_qa_ratio(all)": hallucination / total,
"omission_qa_ratio(all)": omission / total,
"qa_valid_num": valid,
"qa_num": total
}
if valid > 0:
metrics.update({
"correct_qa_ratio(valid)": correct / valid,
"hallucination_qa_ratio(valid)": hallucination / valid,
"omission_qa_ratio(valid)": omission / valid
})
else:
metrics.update({
"correct_qa_ratio(valid)": 0,
"hallucination_qa_ratio(valid)": 0,
"omission_qa_ratio(valid)": 0
})
return metrics
@staticmethod
def compute_time_metrics(eval_results_file: str) -> dict[str, float]:
"""Compute timing metrics from evaluation results."""
answer_duration = 0
with open(eval_results_file, "r", encoding="utf-8") as f:
for line in f:
if not line.strip():
continue
user_data = json.loads(line)
for session in user_data["sessions"]:
eval_results = session.get("evaluation_results", {})
for qa in eval_results.get("question_answering_records", []):
answer_duration += qa.get("answer_duration_ms", 0)
# Convert to minutes
return {
"answer_duration_time": answer_duration / 1000 / 60,
"total_duration_time": answer_duration / 1000 / 60
}
# ==================== Main Pipeline ====================
class HaluMemBaselineEvaluator:
"""Main evaluator orchestrating the baseline evaluation pipeline."""
def __init__(self, config: EvalConfig):
self.config = config
self.file_manager = FileManager(config.output_dir)
self.qa_evaluator = BaselineQuestionAnsweringEvaluator()
self.data_loader = DataLoader()
async def process_session(
self,
session: dict,
session_id: int,
user_name: str,
uuid: str
) -> dict:
"""Process a single session."""
session_data = {
"uuid": uuid,
"user_name": user_name,
"session_id": session_id,
"memory_points": session["memory_points"]
}
# Skip generated QA sessions
if session.get("is_generated_qa_session", False):
session_data["is_generated_qa_session"] = True
return session_data
# Store dialogue
dialogue = session["dialogue"]
session_data["dialogue"] = dialogue
# Evaluate questions if present
if "questions" in session:
formatted_dialogue = self.data_loader.format_dialogue_for_eval(dialogue, user_name)
qa_results = await self.qa_evaluator.evaluate_questions(
questions=session["questions"],
user_name=user_name,
uuid=uuid,
session_id=session_id,
formatted_dialogue=formatted_dialogue
)
session_data["evaluation_results"] = {
"question_answering_records": qa_results
}
return session_data
async def process_user(self, user_data: dict) -> dict:
"""Process all sessions for a user."""
user_name = self.data_loader.extract_user_name(user_data["persona_info"])
uuid = user_data["uuid"]
logger.info(f"Processing user: {user_name}")
for idx, session in enumerate(user_data["sessions"]):
logger.info(f" Session {idx + 1}/{len(user_data['sessions'])}")
session_data = await self.process_session(
session=session,
session_id=idx,
user_name=user_name,
uuid=uuid
)
self.file_manager.save_session(user_name, idx, session_data)
return {"uuid": uuid, "user_name": user_name, "status": "ok"}
async def run_evaluation(self):
"""Run the complete evaluation pipeline."""
start_time = time.time()
# Load user data
all_users = self.data_loader.load_jsonl(self.config.data_path)
users_to_process = all_users[:self.config.user_num]
print("\n" + "=" * 80)
print("HALUMEM BASELINE EVALUATION - DIRECT QA WITHOUT MEMORY SYSTEM")
print(f"Users: {len(users_to_process)} | Concurrency: {self.config.max_concurrency}")
print("=" * 80 + "\n")
# Process users with concurrency control
semaphore = asyncio.Semaphore(self.config.max_concurrency)
async def process_with_cache_check(idx: int, user_data: dict):
async with semaphore:
user_name = self.data_loader.extract_user_name(user_data["persona_info"])
# Check cache
if self.file_manager.user_has_cache(user_name):
print(f"⚡ [{idx}/{len(users_to_process)}] Skipping {user_name} (cached)")
return {"user_name": user_name, "status": "cached"}
print(f"🔄 [{idx}/{len(users_to_process)}] Processing {user_name}...")
result = await self.process_user(user_data)
print(f"✅ [{idx}/{len(users_to_process)}] Completed {user_name}")
return result
tasks = [
process_with_cache_check(idx, user)
for idx, user in enumerate(users_to_process, 1)
]
await asyncio.gather(*tasks)
# Combine results
output_file = os.path.join(self.config.output_dir, "eval_results.jsonl")
self.file_manager.combine_results(output_file)
elapsed = time.time() - start_time
print(f"\n✅ Processing completed in {elapsed:.2f}s")
print(f"📁 Results: {output_file}\n")
# Aggregate metrics
await self.aggregate_and_report(output_file)
async def aggregate_and_report(self, results_file: str):
"""Aggregate results and generate final report."""
print("=" * 80)
print("AGGREGATING METRICS")
print("=" * 80 + "\n")
# Collect all QA records
qa_records = []
with open(results_file, "r", encoding="utf-8") as f:
for line in f:
if not line.strip():
continue
user_data = json.loads(line)
for session in user_data["sessions"]:
if session.get("is_generated_qa_session"):
continue
eval_results = session.get("evaluation_results", {})
qa_records.extend(
eval_results.get("question_answering_records", [])
)
# Compute metrics
qa_metrics = MetricsAggregator.compute_qa_metrics(qa_records)
time_metrics = MetricsAggregator.compute_time_metrics(results_file)
final_results = {
"overall_score": {
"question_answering": qa_metrics,
"time_consuming": time_metrics
},
"question_answering_records": qa_records
}
# Save final report
report_file = os.path.join(self.config.output_dir, "eval_statistics.json")
with open(report_file, "w", encoding="utf-8") as f:
json.dump(final_results, f, ensure_ascii=False, indent=4)
print(f"📊 Statistics saved to: {report_file}\n")
# Print summary
self._print_summary(qa_metrics, time_metrics)
def _print_summary(self, qa_metrics: dict, time_metrics: dict):
"""Print evaluation summary."""
print("=" * 80)
print("EVALUATION SUMMARY")
print("=" * 80 + "\n")
print("📊 Question Answering:")
print(f" Correct (all): {qa_metrics['correct_qa_ratio(all)']:.4f}")
print(f" Hallucination (all): {qa_metrics['hallucination_qa_ratio(all)']:.4f}")
print(f" Omission (all): {qa_metrics['omission_qa_ratio(all)']:.4f}")
print(f" Correct (valid): {qa_metrics['correct_qa_ratio(valid)']:.4f}")
print(f" Hallucination (valid): {qa_metrics['hallucination_qa_ratio(valid)']:.4f}")
print(f" Omission (valid): {qa_metrics['omission_qa_ratio(valid)']:.4f}")
print(f" Valid/Total: {qa_metrics['qa_valid_num']}/{qa_metrics['qa_num']}")
print(f"\n⏱️ Time Metrics:")
print(f" Answer Duration: {time_metrics['answer_duration_time']:.2f} min")
print(f" Total: {time_metrics['total_duration_time']:.2f} min")
print("\n" + "=" * 80)
# ==================== Entry Point ====================
def main(
data_path: str,
user_num: int = 1,
max_concurrency: int = 2
):
"""Main entry point."""
config = EvalConfig(
data_path=data_path,
user_num=user_num,
max_concurrency=max_concurrency
)
evaluator = HaluMemBaselineEvaluator(config)
asyncio.run(evaluator.run_evaluation())
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(
description="Evaluate Baseline (Direct QA) on HaluMem benchmark"
)
parser.add_argument(
"--data_path",
type=str,
required=True,
help="Path to HaluMem JSONL file"
)
parser.add_argument(
"--user_num",
type=int,
default=1,
help="Number of users to evaluate (default: 1)"
)
parser.add_argument(
"--max_concurrency",
type=int,
default=2,
help="Maximum concurrent user processing (default: 2)"
)
args = parser.parse_args()
main(
data_path=args.data_path,
user_num=args.user_num,
max_concurrency=args.max_concurrency
)

File diff suppressed because it is too large Load diff

View file

@ -142,6 +142,9 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
if op.memory_nodes:
self.memory_nodes.extend(op.memory_nodes)
if hasattr(op, "messages") and op.messages:
self.messages.extend(op.messages)
tool_result = str(op.output)
tool_message = Message(
role=Role.TOOL,

View file

@ -9,6 +9,8 @@ from ...core.utils import format_messages
@C.register_op()
class PersonalSummarizerV2(BaseMemoryAgent):
memory_type: MemoryType = MemoryType.PERSONAL
"""Simplified personal memory summarizer that uses v2 memory tools.
This summarizer follows a three-step workflow:
@ -17,11 +19,6 @@ class PersonalSummarizerV2(BaseMemoryAgent):
3. UpdateMemories: Delete outdated memories and add new ones
"""
def __init__(self, **kwargs):
super().__init__(**kwargs)
memory_type: MemoryType = MemoryType.PERSONAL
def _build_tool_call(self) -> ToolCall:
"""Build tool call schema for the agent."""
return ToolCall(
@ -83,28 +80,28 @@ class PersonalSummarizerV2(BaseMemoryAgent):
**kwargs,
)
# Check if AddMemoryDrafts tool was executed
exist_memory_drafts = False
if assistant_message.tool_calls:
for tool_call in assistant_message.tool_calls:
if tool_call.name == "add_memory_drafts":
exist_memory_drafts = True
break
# If memory drafts were added, regenerate system prompt with simplified context
if exist_memory_drafts:
simplified_context = "The conversation context has been summarized in memory drafts."
new_system_prompt = self.prompt_format(
prompt_name="system_prompt",
context=simplified_context,
memory_type=self.memory_type.value,
memory_target=self.memory_target,
)
# Update the system message in the message history
for i, msg in enumerate(self.messages):
if msg.role == Role.SYSTEM:
self.messages[i] = Message(role=Role.SYSTEM, content=new_system_prompt)
break
# # Check if AddMemoryDrafts tool was executed
# exist_memory_drafts = False
# if assistant_message.tool_calls:
# for tool_call in assistant_message.tool_calls:
# if tool_call.name == "add_memory_drafts":
# exist_memory_drafts = True
# break
#
# # If memory drafts were added, regenerate system prompt with simplified context
# if exist_memory_drafts:
# simplified_context = "The conversation context has been summarized in memory drafts."
# new_system_prompt = self.prompt_format(
# prompt_name="system_prompt",
# context=simplified_context,
# memory_type=self.memory_type.value,
# memory_target=self.memory_target,
# )
#
# # Update the system message in the message history
# for i, msg in enumerate(self.messages):
# if msg.role == Role.SYSTEM:
# self.messages[i] = Message(role=Role.SYSTEM, content=new_system_prompt)
# break
return messages

View file

@ -3,58 +3,38 @@ tool: |
Use this tool to analyze dialogues and extract important personal information about users,
such as preferences, habits, personal background, relationships, and significant facts.
# - **Memory granularity**: Each memory should record ONE complete piece of information - don't pack multiple facts into one memory, and don't split a single fact into multiple memories.
# - **Self-contained**: Each memory entry must be self-contained and understandable without additional context.
system_prompt: |
You are a professional memory agent. Your task is to update the main agent's {memory_type} memory regarding {memory_target} based on the context.
You are a professional memory agent managing **{memory_type}** memories about **{memory_target}** for the main agent.
**CRITICAL**: You must extract and store information STRICTLY based on what is explicitly stated in the context. DO NOT infer, assume, fabricate, or add any information that is not directly present in the dialogue. Only extract facts that are clearly and explicitly mentioned.
## Context:
## Latest Conversation:
The context below contains the most recent conversation. Each message is formatted as: `round<index> [<timestamp>] <role/name>: <content>` where timestamp is `YYYY-MM-DD HH:MM:SS`.
{context}
**Context Format Explanation**:
The context contains formatted conversation messages in the following structure:
- Each message is formatted as: `round<index> [<timestamp>] <role/name>: <content>`
- The timestamp is in format: `YYYY-MM-DD HH:MM:SS`
- Content may include reasoning, tool calls
- **Time metadata handling**: When extracting memories with time information, store year/month/day in the metadata. For relative time references (e.g., "last year", "two months ago"), calculate the actual date based on the message's timestamp and store the calculated year/month/day in metadata
## Memory Objective:
You are managing **{memory_type}** memories about **{memory_target}** for the main agent. Focus on extracting and storing information directly related to this person's preferences, habits, personal background, and significant facts.
**CRITICAL**: Extract information ONLY from what is explicitly stated. DO NOT infer, assume, or fabricate any information.
## Your Tasks - Three-Step Workflow:
## Your Tasks
### Step 1: Generate Memory Drafts
Use the `AddMemoryDrafts` tool to create initial memory drafts from the conversation context.
- **Analyze the context**: Determine whether the conversation contains important, memorable information, including but not limited to: user preferences, habits, or personal details; key facts, decisions, or conclusions; relationships or contextual background related to people or topics.
- **Extract key information**: Create memory drafts using clear and concise phrasing **strictly based on what is explicitly stated in the context**.
- **Important**: DO NOT infer, assume, or add any information beyond what is directly mentioned in the conversation.
- **Time references**: If the context involves time references (like "last year", "two months ago", etc.), you **must** calculate the actual date based on the context's timestamp metadata. For example, if a memory from May 4, 2022 mentions "went to India last year," then the trip occurred in 2021. Include this calculated time information in the memory metadata (year, month, day).
- **Memory granularity**: Each memory should record ONE complete piece of information - don't pack multiple facts into one memory, and don't split a single fact into multiple memories.
- **Self-contained**: Each memory entry must be self-contained and understandable without additional context.
Use `AddMemoryDrafts` to extract key facts from the latest conversation.
- Extract important information: preferences, habits, currentstatus, personal details, key facts, decisions, or conclusions.
- Use clear, concise phrasing based strictly on explicit statements.
- Record the timestamp of the source message for each memory including the year, month, and day.
### Step 2: Retrieve Similar and Recent Memories
Use the `RetrieveRecentAndSimilarMemories` tool to find existing related memories.
- **For EACH memory draft**, perform a semantic similarity search to find existing, potentially relevant memories.
- **Example**: For "Person A was born on date X", search for "Person A birth date age".
- **Retrieve comprehensively**: Retrieve all related memories for thorough comparison to prevent any duplication or conflicts.
Use `RetrieveRecentAndSimilarMemories` to query historical memories for each draft.
- Search for semantically similar memories and recent memories.
- This ensures Step 3 avoids duplicates and properly updates existing memories.
### Step 3: Update Memories
Use the `UpdateMemories` tool to finalize the memory updates.
- **Compare and decide**: Compare the newly extracted memory drafts with the retrieved memories from Step 2.
- **CRITICAL DEDUPLICATION CHECK**: Before adding ANY new memory:
- Check if the SAME INFORMATION already exists in retrieved memories
- Consider memories as duplicates even if wording differs, as long as they convey the SAME core fact
- Examples of duplicate information:
* "Person A was born on date X. He/She is N years old." vs "Person A is a gender born on date X. He/She is currently N years old." → DUPLICATES
* "Lives in city" vs "Person A lives in city" → DUPLICATES
* "Holds a Bachelor's degree in field" vs "Person A holds a Bachelor's degree in field" → DUPLICATES
- **Choose the appropriate operation**:
- **If the information already exists and is consistent**: SKIP—fill empty array in `memory_ids_to_delete` and `memories_to_add`. Do NOT add duplicate memories.
- **If existing memory needs supplementation with NEW details**: Delete the old memory (add its ID to `memory_ids_to_delete`), then add the enhanced consolidated version to `memories_to_add`.
- **If existing memory is outdated or contradicted**: Delete it (add ID to `memory_ids_to_delete`), then add the corrected version to `memories_to_add`.
- **If multiple memories contain similar/overlapping information**: Delete all duplicates (add IDs to `memory_ids_to_delete`), then add one merged memory to `memories_to_add`.
- **If the information is entirely new**: Fill empty array in `memory_ids_to_delete`, and add the new memory to `memories_to_add`.
Use `UpdateMemories` to update the memory store by combining drafts with historical memories.
- **Delete conflicts**: Remove old memories that contradict the new drafts (keep most recent/accurate).
- **Add new**: Add drafts that represent completely new information.
- **Skip duplicates**: Do not add drafts that duplicate existing memories.
- **Preserve others**: Keep unrelated historical memories unchanged.
- Write concise memories using minimum words needed. Ensure no information loss.
user_message: |
Please analyze the context and update the memory store following the three-step workflow:

View file

@ -21,7 +21,7 @@ class ReMeSummarizerV2(BaseMemoryAgent):
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": self.prompt_format("tool"),
"description": self.get_prompt("tool"),
"parameters": {
"type": "object",
"properties": {

View file

@ -118,7 +118,7 @@ class BaseMemoryTool(BaseOp, metaclass=ABCMeta):
metadata=metadata or {},
)
logger.opt(depth=1).info(
f"[{self.__class__.__name__}] build node={node.model_dump_json(indent=2, exclude_none=True)}",
)
# logger.opt(depth=1).info(
# f"[{self.__class__.__name__}] build node={node.model_dump_json(indent=2, exclude_none=True)}",
# )
return node

View file

@ -102,29 +102,5 @@ class AddMemoryDrafts(BaseMemoryTool):
async def execute(self):
"""Execute add drafts operation: create memory drafts without persisting to vector store."""
# Get memory drafts to add
memory_drafts = self.context.get("memory_drafts", [])
# Validate input
if not memory_drafts:
self.output = "No memory drafts provided. Please provide at least one draft memory."
return
# Build memory nodes (without persisting)
memory_nodes = []
for mem in memory_drafts:
memory_content, when_to_use, metadata = self._extract_memory_data(mem)
if not memory_content:
logger.warning("Skipping memory draft with empty content")
continue
memory_nodes.append(self._build_memory_node(memory_content, when_to_use=when_to_use, metadata=metadata))
if memory_nodes:
self.memory_nodes.extend(memory_nodes)
draft_count = len(memory_nodes)
self.output = f"Successfully created {draft_count} memory draft(s). These drafts are not yet persisted to the vector store."
logger.info(self.output)
else:
self.output = "No valid memory drafts created. Please check your input."
logger.warning(self.output)
self.output = f"Successfully created memory draft(s). These drafts are not yet persisted to the vector store."
logger.info(self.output)

View file

@ -161,9 +161,6 @@ class RetrieveRecentAndSimilarMemories(BaseMemoryTool):
# Update retrieved_nodes in context with new memories
self.retrieved_nodes.extend(new_memory_nodes)
# Set output to new memories only (after deduplication)
self.memory_nodes = new_memory_nodes
if not new_memory_nodes:
self.output = "No new memory_nodes found (duplicates removed)."
else:

View file

@ -8,7 +8,7 @@ from loguru import logger
from ..base_memory_tool import BaseMemoryTool
from ...core.context import C
from ...core.enumeration import MemoryType
from ...core.schema import MemoryNode
from ...core.schema import MemoryNode, Message
if TYPE_CHECKING:
from ...mem_agent import BaseMemoryAgent
@ -26,6 +26,7 @@ class SummaryAndHandsOff(BaseMemoryTool):
from ...mem_agent import BaseMemoryAgent
self.sub_ops: list[BaseMemoryAgent] = [a for a in self.sub_ops if isinstance(a, BaseMemoryAgent)]
self.messages: list[Message] = []
@property
def memory_agent_dict(self) -> dict[MemoryType, "BaseMemoryAgent"]:
@ -145,6 +146,9 @@ class SummaryAndHandsOff(BaseMemoryTool):
if agent.memory_nodes:
self.memory_nodes.extend(agent.memory_nodes)
if agent.messages:
self.messages.extend(agent.messages)
results.append(
{
"memory_type": memory_type.value,

View file

@ -111,13 +111,15 @@ class UpdateMemories(BaseMemoryTool):
# Get removal IDs
memory_ids_to_delete = self.context.get("memory_ids_to_delete", [])
memory_ids_to_delete = [m for m in memory_ids_to_delete if m]
# Deduplicate memory IDs to avoid redundant deletions
memory_ids_to_delete = list(dict.fromkeys(memory_ids_to_delete))
# Get memories to add
memories_to_add = self.context.get("memories_to_add", [])
# Validate input
if not memory_ids_to_delete and not memories_to_add:
self.output = "No memories to remove or add. Please provide at least one operation."
self.output = "No memories to remove or add. Operation has been done."
return
removed_count = 0
@ -164,6 +166,6 @@ class UpdateMemories(BaseMemoryTool):
if operations:
self.output = f"Successfully {' and '.join(operations)} in vector_store."
else:
self.output = "No valid operations performed. Please check your input."
self.output = "Operation has been done."
logger.info(self.output)

View file

@ -11,6 +11,7 @@ memory_ids_to_delete: |
A list of unique identifiers (memory_ids) of the memories to remove.
Each ID should be a valid memory_id obtained from previous memory retrieval or addition operations.
These memories will be removed before adding the new updated memories.
**IMPORTANT**: Do NOT add duplicate memory_ids. Each memory_id should appear only once in the list.
memories_to_add: |
A list of new memory objects to add after removal.

View file

@ -219,9 +219,9 @@ class ReMe(Application):
if user_id:
metadata_desc = {
"year": "The year when the memory content occurred.",
"month": "The month when the memory content occurred.",
"day": "The day when the memory content occurred.",
"year": "The year when the message content occurred.",
"month": "The month when the message content occurred.",
"day": "The day when the message content occurred.",
}
meta_memories = [
{
@ -256,12 +256,12 @@ class ReMe(Application):
],
)
try:
await reme_summarizer_v2.call(messages=messages, description=description, **kwargs)
return personal_summarizer_v2.memory_nodes, personal_summarizer_v2.messages, personal_summarizer_v2.success
except Exception as e:
print(f"Warning: reme_summarizer_v2.call failed: {e}")
return [], [], False
# try:
await reme_summarizer_v2.call(messages=messages, description=description, **kwargs)
return reme_summarizer_v2.memory_nodes, reme_summarizer_v2.messages, reme_summarizer_v2.success
# except Exception as e:
# print(f"Warning: reme_summarizer_v2.call failed: {e}")
# return [], [], False
else:
raise NotImplementedError
@ -305,12 +305,12 @@ class ReMe(Application):
],
)
try:
await reme_retriever_v2.call(query=query, messages=messages, description=description, **kwargs)
return reme_retriever_v2.output, reme_retriever_v2.messages, reme_retriever_v2.success
except Exception as e:
print(f"Warning: reme_retriever_v2.call failed: {e}")
return "error, not retrieved", [], False
# try:
await reme_retriever_v2.call(query=query, messages=messages, description=description, **kwargs)
return reme_retriever_v2.output, reme_retriever_v2.messages, reme_retriever_v2.success
# except Exception as e:
# print(f"Warning: reme_retriever_v2.call failed: {e}")
# return "error, not retrieved", [], False
else:
raise NotImplementedError