mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-08-28 05:25:04 +00:00
Add bool trigger on Profile memory ; Update Longmemeval Eval and HalumemEval (#129)
* feat(reme): 添加配置选项以启用或禁用个人资料功能
This commit is contained in:
parent
ab4dbd6220
commit
9294e65dfb
3 changed files with 255 additions and 183 deletions
|
|
@ -191,7 +191,7 @@ async def answer_question_with_memories(
|
|||
question: str,
|
||||
memories: str,
|
||||
user_id: str = None,
|
||||
model_name: str = "qwen3-30b-a3b-instruct-2507",
|
||||
eval_model_name: str = "qwen3-30b-a3b-instruct-2507",
|
||||
):
|
||||
"""
|
||||
Answer a question using retrieved memories with PROMPT_MEMZERO_JSON template.
|
||||
|
|
@ -201,7 +201,7 @@ async def answer_question_with_memories(
|
|||
question: The question to answer
|
||||
memories: The retrieved memories (formatted as context)
|
||||
user_id: Optional user ID for context formatting
|
||||
model_name: Model name to use for LLM request
|
||||
eval_model_name: Model name to use for LLM request
|
||||
|
||||
Returns:
|
||||
dict with 'reasoning' and 'answer' fields
|
||||
|
|
@ -223,7 +223,7 @@ async def answer_question_with_memories(
|
|||
question=question,
|
||||
)
|
||||
|
||||
result = await reme.get_llm(model_name).simple_request_for_json(
|
||||
result = await reme.get_llm(eval_model_name).simple_request_for_json(
|
||||
prompt=prompt,
|
||||
model_name=None,
|
||||
)
|
||||
|
|
@ -236,6 +236,7 @@ async def evaluation_for_memory_accuracy(
|
|||
dialogue: str,
|
||||
golden_memories: list[dict],
|
||||
candidate_memory: dict,
|
||||
eval_model_name: str = "qwen-flash",
|
||||
):
|
||||
"""
|
||||
Memory Accuracy Evaluation - Check if an extracted memory is accurate.
|
||||
|
|
@ -245,7 +246,7 @@ async def evaluation_for_memory_accuracy(
|
|||
dialogue: The formatted dialogue string
|
||||
golden_memories: List of golden memory points from the session
|
||||
candidate_memory: The extracted memory to evaluate
|
||||
model_name: Model name to use for LLM request
|
||||
eval_model_name: Model name to use for LLM request
|
||||
|
||||
Returns:
|
||||
dict with 'accuracy_score' (0/1/2), 'is_included_in_golden_memories' (true/false), and 'reason'
|
||||
|
|
@ -265,7 +266,7 @@ async def evaluation_for_memory_accuracy(
|
|||
candidate_memory=candidate_content,
|
||||
)
|
||||
|
||||
result = await reme.get_llm("qwen-flash").simple_request_for_json(
|
||||
result = await reme.get_llm(eval_model_name).simple_request_for_json(
|
||||
prompt=prompt,
|
||||
model_name=None,
|
||||
)
|
||||
|
|
@ -454,7 +455,7 @@ class MemoryProcessor:
|
|||
question=query,
|
||||
memories=memories,
|
||||
user_id=user_id,
|
||||
model_name=self.eval_model_name,
|
||||
eval_model_name=self.eval_model_name,
|
||||
)
|
||||
|
||||
# Add original memories to the result
|
||||
|
|
@ -1420,8 +1421,7 @@ if __name__ == "__main__":
|
|||
parser.add_argument(
|
||||
"--data_path",
|
||||
type=str,
|
||||
# required=True,
|
||||
default="/Users/zhouwk/PycharmProjects/MemAgent/dataset/halumem/HaluMem-Medium.jsonl",
|
||||
required=True,
|
||||
help="Path to HaluMem JSONL file",
|
||||
)
|
||||
parser.add_argument(
|
||||
|
|
|
|||
|
|
@ -16,7 +16,6 @@ Usage:
|
|||
|
||||
import asyncio
|
||||
import json
|
||||
import shutil
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone, timedelta
|
||||
|
|
@ -42,8 +41,9 @@ class EvalConfig:
|
|||
max_concurrency: int = 1
|
||||
batch_size: int = 30
|
||||
output_dir: str = "cache/bench_results/longmemeval_reme"
|
||||
reme_model_name: str = "qwen-flash"
|
||||
eval_model_name: str = "qwen3-max"
|
||||
reme_model_name: str = "qwen-flash" # summary模型
|
||||
retrieve_model_name: str = "qwen-max" # retrieve模型
|
||||
eval_model_name: str = "qwen-max" # 评估/判断模型
|
||||
algo_version: str = "v1"
|
||||
samples_per_type: int = -1 # Number of samples per question type, -1 for all
|
||||
enable_thinking_params: bool = False
|
||||
|
|
@ -254,7 +254,7 @@ async def answer_question_with_memories(
|
|||
question: str,
|
||||
memories: str,
|
||||
user_id: str = None,
|
||||
model_name: str = "qwen3-max",
|
||||
model_name: str = "qwen-max",
|
||||
):
|
||||
"""
|
||||
Answer a question using retrieved memories with PROMPT_MEMZERO_JSON template.
|
||||
|
|
@ -286,9 +286,9 @@ async def answer_question_with_memories(
|
|||
question=question,
|
||||
)
|
||||
|
||||
result = await reme.get_llm(model_name).simple_request_for_json(
|
||||
result = await reme.default_llm.simple_request_for_json(
|
||||
prompt=prompt,
|
||||
model_name=None,
|
||||
model_name=model_name,
|
||||
)
|
||||
|
||||
return result
|
||||
|
|
@ -304,12 +304,14 @@ class MemoryProcessor:
|
|||
self,
|
||||
reme: ReMe,
|
||||
reme_model_name: str = "qwen-flash",
|
||||
eval_model_name: str = "qwen3-max",
|
||||
retrieve_model_name: str = "qwen-max",
|
||||
eval_model_name: str = "qwen-max",
|
||||
algo_version: str = "v1",
|
||||
enable_thinking_params: bool = False,
|
||||
):
|
||||
self.reme = reme
|
||||
self.reme_model_name = reme_model_name
|
||||
self.retrieve_model_name = retrieve_model_name
|
||||
self.eval_model_name = eval_model_name
|
||||
self.algo_version = algo_version
|
||||
self.enable_thinking_params = enable_thinking_params
|
||||
|
|
@ -369,7 +371,7 @@ class MemoryProcessor:
|
|||
|
||||
# Retrieve memories from ReMe using new API
|
||||
result = await self.reme.retrieve_memory(
|
||||
llm_config_name="qwen3-max-think",
|
||||
llm_config_name=self.retrieve_model_name,
|
||||
query=query,
|
||||
retrieve_top_k=top_k,
|
||||
user_name=user_id,
|
||||
|
|
@ -541,63 +543,84 @@ class LongMemEvalEvaluator:
|
|||
|
||||
def __init__(self, config: EvalConfig):
|
||||
self.config = config
|
||||
self.reme = ReMe(
|
||||
self.file_manager = FileManager(config.output_dir)
|
||||
self.data_loader = DataLoader()
|
||||
|
||||
# Store LLM configs for creating ReMe instances per question
|
||||
self._llm_configs = {
|
||||
"qwen-plus-t": {
|
||||
"backend": "openai",
|
||||
"model_name": "qwen-plus",
|
||||
"extra_body": {
|
||||
"enable_thinking": True,
|
||||
},
|
||||
},
|
||||
"qwen-max-t": {
|
||||
"backend": "openai",
|
||||
"model_name": "qwen3-max",
|
||||
"extra_body": {
|
||||
"enable_thinking": True,
|
||||
},
|
||||
},
|
||||
"gpt-4o-mini": {
|
||||
"backend": "openai",
|
||||
"model_name": "gpt-4o-mini-2024-07-18",
|
||||
},
|
||||
"gpt-4o-mini-2024-07-18": {
|
||||
"backend": "openai",
|
||||
"model_name": "gpt-4o-mini-2024-07-18",
|
||||
},
|
||||
"qwen-flash": {
|
||||
"backend": "openai",
|
||||
"model_name": "qwen-flash",
|
||||
},
|
||||
"qwen-max": {
|
||||
"backend": "openai",
|
||||
"model_name": "qwen3-max",
|
||||
},
|
||||
}
|
||||
|
||||
# Load evaluation prompts path
|
||||
self._prompts_yaml_path = Path(__file__).parent / "eval_reme.yaml"
|
||||
|
||||
def _create_reme_for_question(self, question_id: str) -> ReMe:
|
||||
"""Create a ReMe instance for a specific question with isolated collection.
|
||||
|
||||
Args:
|
||||
question_id: The question ID to use as collection name
|
||||
|
||||
Returns:
|
||||
ReMe instance with isolated vector store collection
|
||||
"""
|
||||
collection_name = f"longmemeval_{question_id}"
|
||||
reme = ReMe(
|
||||
default_llm_config={
|
||||
"model_name": self.config.reme_model_name,
|
||||
},
|
||||
llms={
|
||||
"qwen-plus-think": {
|
||||
"backend": "openai",
|
||||
"model_name": "qwen-plus",
|
||||
"extra_body": {
|
||||
"enable_thinking": True,
|
||||
},
|
||||
},
|
||||
"qwen3-max-think": {
|
||||
"backend": "openai",
|
||||
"model_name": "qwen3-max",
|
||||
"extra_body": {
|
||||
"enable_thinking": True,
|
||||
},
|
||||
},
|
||||
"qwen3-max": {
|
||||
"backend": "openai",
|
||||
"model_name": "qwen3-max",
|
||||
"extra_body": {
|
||||
"enable_thinking": False,
|
||||
},
|
||||
},
|
||||
default_vector_store_config={
|
||||
"collection_name": collection_name,
|
||||
},
|
||||
llms=self._llm_configs,
|
||||
)
|
||||
|
||||
# Load evaluation prompts into ReMe's prompt handler
|
||||
prompts_yaml_path = Path(__file__).parent / "eval_reme.yaml"
|
||||
self.reme.prompt_handler.load_prompt_by_file(prompts_yaml_path)
|
||||
# Load evaluation prompts
|
||||
reme.prompt_handler.load_prompt_by_file(self._prompts_yaml_path)
|
||||
|
||||
self.file_manager = FileManager(config.output_dir)
|
||||
self.memory_processor = MemoryProcessor(
|
||||
self.reme,
|
||||
config.reme_model_name,
|
||||
config.eval_model_name,
|
||||
config.algo_version,
|
||||
config.enable_thinking_params,
|
||||
)
|
||||
self.judge = LongMemEvalJudge(self.reme, config.eval_model_name)
|
||||
self.data_loader = DataLoader()
|
||||
return reme
|
||||
|
||||
async def __aenter__(self):
|
||||
"""Async context manager entry."""
|
||||
await self.reme.start()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||||
"""Async context manager exit with cleanup."""
|
||||
await self.reme.close()
|
||||
return False
|
||||
|
||||
async def process_question_entry(self, entry: dict, idx: int) -> dict:
|
||||
"""Process a single question entry.
|
||||
|
||||
Each question gets its own ReMe instance with isolated vector store collection.
|
||||
|
||||
Args:
|
||||
entry: A question entry from LongMemEval dataset
|
||||
idx: Index of the question
|
||||
|
|
@ -614,8 +637,8 @@ class LongMemEvalEvaluator:
|
|||
haystack_session_ids = entry["haystack_session_ids"]
|
||||
haystack_sessions = entry["haystack_sessions"]
|
||||
|
||||
# Use question_id as user_id for isolation
|
||||
user_id = f"longmemeval_{question_id}"
|
||||
# Use "User" as user_name, question_id is stored in collection_name
|
||||
user_name = "User"
|
||||
|
||||
logger.info(f"\n{'=' * 60}")
|
||||
logger.info(f"Question ID: {question_id}")
|
||||
|
|
@ -626,103 +649,116 @@ class LongMemEvalEvaluator:
|
|||
logger.info(f"Number of sessions: {len(haystack_sessions)}")
|
||||
logger.info(f"{'=' * 60}")
|
||||
|
||||
# Step 2: Process all haystack sessions to build memory
|
||||
all_extracted_memories = []
|
||||
all_agent_messages = []
|
||||
total_summary_duration_ms = 0
|
||||
# Create isolated ReMe instance for this question
|
||||
reme = self._create_reme_for_question(question_id)
|
||||
await reme.start()
|
||||
|
||||
for session_idx, (session, session_date, session_id) in enumerate(
|
||||
zip(haystack_sessions, haystack_dates, haystack_session_ids),
|
||||
):
|
||||
logger.info(f" Processing session {session_idx + 1}/{len(haystack_sessions)}: {session_id}")
|
||||
try:
|
||||
# Create memory processor and judge for this ReMe instance
|
||||
memory_processor = MemoryProcessor(
|
||||
reme,
|
||||
self.config.reme_model_name,
|
||||
self.config.retrieve_model_name,
|
||||
self.config.eval_model_name,
|
||||
self.config.algo_version,
|
||||
self.config.enable_thinking_params,
|
||||
)
|
||||
judge = LongMemEvalJudge(reme, self.config.eval_model_name)
|
||||
|
||||
# Convert session to messages
|
||||
messages = self.data_loader.convert_session_to_messages(session, session_date)
|
||||
# Clear existing vector store data for this collection
|
||||
await reme.default_vector_store.delete_all()
|
||||
|
||||
if not messages:
|
||||
continue
|
||||
# Step 2: Process all haystack sessions to build memory
|
||||
all_extracted_memories = []
|
||||
all_agent_messages = []
|
||||
total_summary_duration_ms = 0
|
||||
|
||||
# Add memories
|
||||
extracted_memories, agent_messages, duration_ms = await self.memory_processor.add_memories(
|
||||
user_id=user_id,
|
||||
messages=messages,
|
||||
batch_size=self.config.batch_size,
|
||||
for session_idx, (session, session_date, session_id) in enumerate(
|
||||
zip(haystack_sessions, haystack_dates, haystack_session_ids),
|
||||
):
|
||||
logger.info(f" Processing session {session_idx + 1}/{len(haystack_sessions)}: {session_id}")
|
||||
|
||||
# Convert session to messages
|
||||
messages = self.data_loader.convert_session_to_messages(session, session_date)
|
||||
|
||||
if not messages:
|
||||
continue
|
||||
|
||||
# Add memories using "User" as user_name
|
||||
extracted_memories, agent_messages, duration_ms = await memory_processor.add_memories(
|
||||
user_id=user_name,
|
||||
messages=messages,
|
||||
batch_size=self.config.batch_size,
|
||||
)
|
||||
|
||||
all_extracted_memories.extend(extracted_memories)
|
||||
all_agent_messages.extend(agent_messages)
|
||||
total_summary_duration_ms += duration_ms
|
||||
|
||||
# Step 3: Search memory and answer question
|
||||
logger.info(" Answering question using ReMe...")
|
||||
answer_dict, retrieve_messages, retrieve_duration_ms = await memory_processor.search_memory(
|
||||
query=f"[Question_date: {question_date} | Question_type: {question_type}] " + question,
|
||||
user_id=user_name,
|
||||
top_k=self.config.top_k,
|
||||
)
|
||||
|
||||
all_extracted_memories.extend(extracted_memories)
|
||||
all_agent_messages.extend(agent_messages)
|
||||
total_summary_duration_ms += duration_ms
|
||||
# Extract answer and reasoning from the structured response
|
||||
model_response = answer_dict.get("answer", "")
|
||||
model_reasoning = answer_dict.get("reasoning", "")
|
||||
retrieved_memories = answer_dict.get("memories", "")
|
||||
retrieved_nodes = answer_dict.get("retrieved_nodes", [])
|
||||
|
||||
# Step 3: Search memory and answer question
|
||||
logger.info(" Answering question using ReMe...")
|
||||
answer_dict, retrieve_messages, retrieve_duration_ms = await self.memory_processor.search_memory(
|
||||
query=f"[Question_date: {question_date} | Question_type: {question_type}] " + question,
|
||||
user_id=user_id,
|
||||
top_k=self.config.top_k,
|
||||
)
|
||||
# Step 4: Judge answer correctness
|
||||
logger.info(" Judging answer correctness...")
|
||||
judgment = await judge.judge_answer(
|
||||
question_type=question_type,
|
||||
question=question,
|
||||
answer=answer,
|
||||
response=model_response,
|
||||
)
|
||||
|
||||
# Extract answer and reasoning from the structured response
|
||||
model_response = answer_dict.get("answer", "")
|
||||
model_reasoning = answer_dict.get("reasoning", "")
|
||||
retrieved_memories = answer_dict.get("memories", "")
|
||||
retrieved_nodes = answer_dict.get("retrieved_nodes", [])
|
||||
is_correct = judgment.get("is_correct")
|
||||
logger.info(
|
||||
f" → Answer judgment: {'Correct' if is_correct else 'Incorrect' if is_correct is False else 'Error'}",
|
||||
)
|
||||
|
||||
# Step 4: Judge answer correctness
|
||||
logger.info(" Judging answer correctness...")
|
||||
judgment = await self.judge.judge_answer(
|
||||
question_type=question_type,
|
||||
question=question,
|
||||
answer=answer,
|
||||
response=model_response,
|
||||
)
|
||||
result = {
|
||||
"question_id": question_id,
|
||||
"question_type": question_type,
|
||||
"question": question,
|
||||
"answer": answer,
|
||||
"question_date": question_date,
|
||||
"haystack_dates": haystack_dates,
|
||||
"haystack_session_ids": haystack_session_ids,
|
||||
"num_sessions": len(haystack_sessions),
|
||||
"model_response": model_response,
|
||||
"model_reasoning": model_reasoning,
|
||||
"retrieved_memories": retrieved_memories,
|
||||
"retrieved_nodes": retrieved_nodes,
|
||||
"judgment": judgment,
|
||||
"extracted_memories": all_extracted_memories,
|
||||
"summary_duration_ms": total_summary_duration_ms,
|
||||
"retrieve_duration_ms": retrieve_duration_ms,
|
||||
"summary_messages": all_agent_messages,
|
||||
"retrieve_messages": retrieve_messages,
|
||||
}
|
||||
|
||||
is_correct = judgment.get("is_correct")
|
||||
logger.info(
|
||||
f" → Answer judgment: {'Correct' if is_correct else 'Incorrect' if is_correct is False else 'Error'}",
|
||||
)
|
||||
# Save individual result
|
||||
self.file_manager.save_question_result(idx, question_id, result)
|
||||
|
||||
result = {
|
||||
"question_id": question_id,
|
||||
"question_type": question_type,
|
||||
"question": question,
|
||||
"answer": answer,
|
||||
"question_date": question_date,
|
||||
"haystack_dates": haystack_dates,
|
||||
"haystack_session_ids": haystack_session_ids,
|
||||
"num_sessions": len(haystack_sessions),
|
||||
"model_response": model_response,
|
||||
"model_reasoning": model_reasoning,
|
||||
"retrieved_memories": retrieved_memories,
|
||||
"retrieved_nodes": retrieved_nodes,
|
||||
"judgment": judgment,
|
||||
"extracted_memories": all_extracted_memories,
|
||||
"summary_duration_ms": total_summary_duration_ms,
|
||||
"retrieve_duration_ms": retrieve_duration_ms,
|
||||
"summary_messages": all_agent_messages,
|
||||
"retrieve_messages": retrieve_messages,
|
||||
}
|
||||
logger.info(f" Question {question_id} - Completed")
|
||||
|
||||
# Save individual result
|
||||
self.file_manager.save_question_result(idx, question_id, result)
|
||||
return result
|
||||
|
||||
logger.info(f" Question {question_id} - Completed")
|
||||
|
||||
return result
|
||||
finally:
|
||||
# Always close the ReMe instance
|
||||
await reme.close()
|
||||
|
||||
async def run_evaluation(self):
|
||||
"""Run the complete evaluation pipeline with parallel processing."""
|
||||
start_time = time.time()
|
||||
|
||||
# Clear existing vector store data
|
||||
await self.reme.default_vector_store.delete_all()
|
||||
|
||||
# Clear meta_memory directory
|
||||
meta_memory_path = Path(f"meta_memory/{self.reme.default_vector_store.collection_name}")
|
||||
if meta_memory_path.exists():
|
||||
shutil.rmtree(meta_memory_path)
|
||||
logger.info(f"Cleared meta_memory directory: {meta_memory_path}")
|
||||
meta_memory_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Load dataset
|
||||
logger.info(f"Loading dataset from: {self.config.data_path}")
|
||||
all_data = self.data_loader.load_json(self.config.data_path)
|
||||
|
|
@ -749,7 +785,7 @@ class LongMemEvalEvaluator:
|
|||
print(f"Samples per type: {self.config.samples_per_type} (-1 = all)")
|
||||
print(f"Questions to process: {total_questions} | Top-K: {self.config.top_k}")
|
||||
print(f"Max Concurrency: {self.config.max_concurrency}")
|
||||
print(f"ReMe Model: {self.config.reme_model_name} | Eval Model: {self.config.eval_model_name}")
|
||||
print(f"Summary Model: {self.config.reme_model_name} | Retrieve Model: {self.config.retrieve_model_name} | Eval Model: {self.config.eval_model_name}")
|
||||
print(f"Algo Version: {self.config.algo_version}")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
|
|
@ -880,7 +916,8 @@ async def main_async(
|
|||
batch_size: int = 30,
|
||||
output_dir: str = "bench_results/longmemeval_reme",
|
||||
reme_model_name: str = "qwen-flash",
|
||||
eval_model_name: str = "qwen3-max",
|
||||
retrieve_model_name: str = "qwen-max",
|
||||
eval_model_name: str = "qwen-max",
|
||||
algo_version: str = "v1",
|
||||
samples_per_type: int = -1,
|
||||
enable_thinking_params: bool = False,
|
||||
|
|
@ -895,6 +932,7 @@ async def main_async(
|
|||
batch_size=batch_size,
|
||||
output_dir=output_dir,
|
||||
reme_model_name=reme_model_name,
|
||||
retrieve_model_name=retrieve_model_name,
|
||||
eval_model_name=eval_model_name,
|
||||
algo_version=algo_version,
|
||||
samples_per_type=samples_per_type,
|
||||
|
|
@ -915,7 +953,8 @@ def main(
|
|||
batch_size: int = 30,
|
||||
output_dir: str = "bench_results/longmemeval_reme",
|
||||
reme_model_name: str = "qwen-flash",
|
||||
eval_model_name: str = "qwen3-max",
|
||||
retrieve_model_name: str = "qwen-max",
|
||||
eval_model_name: str = "qwen-max",
|
||||
algo_version: str = "v1",
|
||||
samples_per_type: int = -1,
|
||||
enable_thinking_params: bool = False,
|
||||
|
|
@ -931,6 +970,7 @@ def main(
|
|||
batch_size=batch_size,
|
||||
output_dir=output_dir,
|
||||
reme_model_name=reme_model_name,
|
||||
retrieve_model_name=retrieve_model_name,
|
||||
eval_model_name=eval_model_name,
|
||||
algo_version=algo_version,
|
||||
samples_per_type=samples_per_type,
|
||||
|
|
@ -948,8 +988,7 @@ if __name__ == "__main__":
|
|||
parser.add_argument(
|
||||
"--data_path",
|
||||
type=str,
|
||||
# default="/Users/zhouwk/PycharmProjects/MemAgent/dataset/longmemeval/longmemeval_s_cleaned.json",
|
||||
default="/Users/zhouwk/PycharmProjects/MemAgent/dataset/longmemeval/longmemeval_oracle.json",
|
||||
required=True,
|
||||
help="Path to LongMemEval JSON file",
|
||||
)
|
||||
parser.add_argument(
|
||||
|
|
@ -979,7 +1018,7 @@ if __name__ == "__main__":
|
|||
parser.add_argument(
|
||||
"--batch_size",
|
||||
type=int,
|
||||
default=10,
|
||||
default=30,
|
||||
help="Batch size for memory summary processing (default: 30)",
|
||||
)
|
||||
parser.add_argument(
|
||||
|
|
@ -991,16 +1030,20 @@ if __name__ == "__main__":
|
|||
parser.add_argument(
|
||||
"--reme_model_name",
|
||||
type=str,
|
||||
# default="gpt-4o-mini-2024-07-18",
|
||||
default="qwen-flash",
|
||||
help="Model name for ReMe operations (default: qwen-flash)",
|
||||
help="Model name for ReMe summary operations (default: qwen-flash)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--retrieve_model_name",
|
||||
type=str,
|
||||
default="qwen-max",
|
||||
help="Model name for memory retrieval (default: qwen-max)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--eval_model_name",
|
||||
type=str,
|
||||
# default="gpt-4o-mini-2024-07-18",
|
||||
default="qwen-max",
|
||||
help="Model name for evaluation/judgment (default: qwen3-max)",
|
||||
default="qwen-flash",
|
||||
help="Model name for evaluation/judgment (default: qwen-max)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--algo_version",
|
||||
|
|
@ -1011,7 +1054,7 @@ if __name__ == "__main__":
|
|||
parser.add_argument(
|
||||
"--samples_per_type",
|
||||
type=int,
|
||||
default=2,
|
||||
default=1,
|
||||
help="Number of samples per question type, -1 for all (default: -1)",
|
||||
)
|
||||
parser.add_argument(
|
||||
|
|
@ -1033,6 +1076,7 @@ if __name__ == "__main__":
|
|||
batch_size=args.batch_size,
|
||||
output_dir=args.output_dir,
|
||||
reme_model_name=args.reme_model_name,
|
||||
retrieve_model_name=args.retrieve_model_name,
|
||||
eval_model_name=args.eval_model_name,
|
||||
algo_version=args.algo_version,
|
||||
samples_per_type=args.samples_per_type,
|
||||
|
|
|
|||
100
reme/reme.py
100
reme/reme.py
|
|
@ -53,6 +53,7 @@ class ReMe(Application):
|
|||
target_user_names: list[str] | None = None,
|
||||
target_task_names: list[str] | None = None,
|
||||
target_tool_names: list[str] | None = None,
|
||||
enable_profile: bool = True,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize ReMe with config.
|
||||
|
|
@ -65,6 +66,10 @@ class ReMe(Application):
|
|||
await reme.retrieve_memory(...)
|
||||
await reme.close()
|
||||
```
|
||||
|
||||
Args:
|
||||
enable_profile: Whether to enable profile functionality. Set to False when using
|
||||
cloud-based vector stores to avoid local file operations. Default is True.
|
||||
"""
|
||||
super().__init__(
|
||||
*args,
|
||||
|
|
@ -83,6 +88,9 @@ class ReMe(Application):
|
|||
default_token_counter_config=default_token_counter_config,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
self.enable_profile = enable_profile
|
||||
|
||||
memory_target_type_mapping: dict[str, MemoryType] = {}
|
||||
if target_user_names:
|
||||
for name in target_user_names:
|
||||
|
|
@ -101,9 +109,12 @@ class ReMe(Application):
|
|||
|
||||
self.service_context.memory_target_type_mapping = memory_target_type_mapping
|
||||
|
||||
profile_path = Path(self.service_context.service_config.working_dir) / "profile"
|
||||
profile_path.mkdir(parents=True, exist_ok=True)
|
||||
self.profile_dir: str = str(profile_path)
|
||||
if self.enable_profile:
|
||||
profile_path = Path(self.service_context.service_config.working_dir) / "profile"
|
||||
profile_path.mkdir(parents=True, exist_ok=True)
|
||||
self.profile_dir: str = str(profile_path)
|
||||
else:
|
||||
self.profile_dir: str = ""
|
||||
|
||||
def _add_meta_memory(self, memory_type: str | MemoryType, memory_target: str):
|
||||
"""Register or validate a memory target with the given memory type."""
|
||||
|
|
@ -174,34 +185,40 @@ class ReMe(Application):
|
|||
format_messages.append(message)
|
||||
|
||||
if version == "default":
|
||||
personal_summarizer_tools = [
|
||||
AddDraftAndRetrieveSimilarMemory(
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
enable_memory_target=False,
|
||||
enable_when_to_use=False,
|
||||
enable_multiple=True,
|
||||
top_k=retrieve_top_k,
|
||||
),
|
||||
AddMemory(
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
enable_memory_target=False,
|
||||
enable_when_to_use=False,
|
||||
enable_multiple=True,
|
||||
),
|
||||
]
|
||||
if self.enable_profile:
|
||||
personal_summarizer_tools.extend(
|
||||
[
|
||||
ReadAllProfiles(
|
||||
enable_thinking_params=False,
|
||||
enable_memory_target=False,
|
||||
profile_dir=self.profile_dir,
|
||||
),
|
||||
UpdateProfilesV1(
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
enable_memory_target=False,
|
||||
enable_multiple=True,
|
||||
profile_dir=self.profile_dir,
|
||||
),
|
||||
],
|
||||
)
|
||||
personal_summarizer: BaseMemoryAgent = PersonalSummarizer(
|
||||
llm=llm_config_name,
|
||||
tools=[
|
||||
AddDraftAndRetrieveSimilarMemory(
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
enable_memory_target=False,
|
||||
enable_when_to_use=False,
|
||||
enable_multiple=True,
|
||||
top_k=retrieve_top_k,
|
||||
),
|
||||
AddMemory(
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
enable_memory_target=False,
|
||||
enable_when_to_use=False,
|
||||
enable_multiple=True,
|
||||
),
|
||||
ReadAllProfiles(
|
||||
enable_thinking_params=False,
|
||||
enable_memory_target=False,
|
||||
profile_dir=self.profile_dir,
|
||||
),
|
||||
UpdateProfilesV1(
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
enable_memory_target=False,
|
||||
enable_multiple=True,
|
||||
profile_dir=self.profile_dir,
|
||||
),
|
||||
],
|
||||
tools=personal_summarizer_tools,
|
||||
)
|
||||
|
||||
else:
|
||||
|
|
@ -324,14 +341,17 @@ class ReMe(Application):
|
|||
"""Retrieve relevant personal, procedural and tool memories for a query."""
|
||||
|
||||
if version == "default":
|
||||
personal_retriever: BaseMemoryAgent = PersonalRetriever(
|
||||
llm=llm_config_name,
|
||||
tools=[
|
||||
personal_retriever_tools = []
|
||||
if self.enable_profile:
|
||||
personal_retriever_tools.append(
|
||||
ReadAllProfiles(
|
||||
enable_thinking_params=False,
|
||||
enable_memory_target=False,
|
||||
profile_dir=self.profile_dir,
|
||||
),
|
||||
)
|
||||
personal_retriever_tools.extend(
|
||||
[
|
||||
RetrieveMemory(
|
||||
top_k=retrieve_top_k,
|
||||
enable_thinking_params=enable_thinking_params,
|
||||
|
|
@ -344,6 +364,10 @@ class ReMe(Application):
|
|||
),
|
||||
],
|
||||
)
|
||||
personal_retriever: BaseMemoryAgent = PersonalRetriever(
|
||||
llm=llm_config_name,
|
||||
tools=personal_retriever_tools,
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(f"version={version} is not supported")
|
||||
|
||||
|
|
@ -601,12 +625,16 @@ class ReMe(Application):
|
|||
return MemoryHandler(memory_target=memory_target, service_context=self.service_context)
|
||||
|
||||
@property
|
||||
def profile_path(self) -> Path:
|
||||
"""Get the path to the profile directory."""
|
||||
def profile_path(self) -> Path | None:
|
||||
"""Get the path to the profile directory. Returns None if profile is disabled."""
|
||||
if not self.enable_profile:
|
||||
return None
|
||||
return Path(self.profile_dir) / self.default_vector_store.collection_name
|
||||
|
||||
def get_profile_handler(self, user_name: str) -> ProfileHandler:
|
||||
"""Get the profile handler for the specified user."""
|
||||
def get_profile_handler(self, user_name: str) -> ProfileHandler | None:
|
||||
"""Get the profile handler for the specified user. Returns None if profile is disabled."""
|
||||
if not self.enable_profile:
|
||||
return None
|
||||
return ProfileHandler(memory_target=user_name, profile_path=self.profile_path)
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue