Merge pull request #97 from nitwtog/main

修正halumem/eval_reme.py评估函数, 添加v2版本用于适配readhistoryv2
This commit is contained in:
jinliyl 2026-02-03 15:56:57 +08:00 • committed by GitHub
commit 1aca24b804
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 83 additions and 28 deletions

View file

@ -38,12 +38,12 @@ class EvalConfig:
data_path: str
top_k: int = 20
user_num: int = 1
max_concurrency: int = 2
batch_size: int = 20
max_concurrency: int = 1
batch_size: int = 40
output_dir: str = "bench_results/reme"
reme_model_name: str = "qwen-flash"
eval_model_name: str = "qwen3-max"
algo_version: str = "halumem"
algo_version: str = "v1"
enable_thinking_params: bool = False
@ -220,7 +220,7 @@ async def answer_question_with_memories(
result = await reme.llm.simple_request_for_json(
prompt=prompt,
model_name="qwen-flash"
model_name=model_name
)
return result
@ -272,8 +272,9 @@ async def evaluation_for_question(
class MemoryProcessor:
"""Handles ReMe memory operations."""
def __init__(self, reme: ReMe, eval_model_name: str = "qwen3-max", algo_version: str = "halumem", enable_thinking_params: bool = False):
def __init__(self, reme: ReMe, reme_model_name:str="qwen3-max",eval_model_name: str = "qwen3-max", algo_version: str = "halumem", enable_thinking_params: bool = False):
self.reme = reme
self.reme_model_name = reme_model_name
self.eval_model_name = eval_model_name
self.algo_version = algo_version
self.enable_thinking_params = enable_thinking_params
@ -353,7 +354,7 @@ class MemoryProcessor:
question=query,
memories=memories,
user_id=user_id,
model_name=self.eval_model_name
model_name=self.eval_model_name,
)
# Add original memories to the result
@ -552,6 +553,7 @@ class HaluMemEvaluator:
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
@ -866,6 +868,7 @@ class HaluMemEvaluator:
async def main_async(
data_path: str,
top_k: int,
batch_size: int,
user_num: int,
max_concurrency: int,
reme_model_name: str= "qwen-flash",
@ -877,6 +880,7 @@ async def main_async(
config = EvalConfig(
data_path=data_path,
top_k=top_k,
batch_size=batch_size,
user_num=user_num,
max_concurrency=max_concurrency,
reme_model_name=reme_model_name,
@ -893,6 +897,7 @@ async def main_async(
def main(
data_path: str,
top_k: int,
batch_size: int,
user_num: int,
max_concurrency: int,
reme_model_name: str= "qwen-flash",
@ -904,6 +909,7 @@ def main(
asyncio.run(main_async(
data_path=data_path,
top_k=top_k,
batch_size=batch_size,
user_num=user_num,
max_concurrency=max_concurrency,
reme_model_name=reme_model_name,
@ -944,6 +950,12 @@ if __name__ == "__main__":
default=1,
help="Maximum concurrent user processing (default: 100)"
)
parser.add_argument(
"--batch_size",
type=int,
default=40,
help="Batch size for memory summary processing of each conversation (default: 40)"
)
parser.add_argument(
"--reme_model_name",
type=str,
@ -959,8 +971,8 @@ if __name__ == "__main__":
parser.add_argument(
"--algo_version",
type=str,
default="halumem",
help="Algorithm version for summary and retrieval (default: halumem)"
default="v1",
help="Algorithm version for summary and retrieval (default: v1)"
)
parser.add_argument(
"--enable_thinking_params",
@ -975,6 +987,7 @@ if __name__ == "__main__":
main(
data_path=args.data_path,
top_k=args.top_k,
batch_size=args.batch_size,
user_num=args.user_num,
max_concurrency=args.max_concurrency,
reme_model_name=args.reme_model_name,

View file

@ -26,6 +26,7 @@ from .tool.memory import (
RetrieveMemory,
DelegateTask,
ReadHistory,
ReadHistoryV2,
ProfileHandler,
MemoryHandler,
AddAndRetrieveSimilarMemory,
@ -238,6 +239,34 @@ class ReMe(Application):
),
],
)
elif version == "v2":
personal_summarizer = PersonalV1Summarizer(
tools=[
AddDraftAndRetrieveSimilarMemory(
enable_thinking_params=enable_thinking_params,
enable_memory_target=False,
enable_when_to_use=False,
enable_multiple=True,
),
UpdateMemoryV1(
enable_thinking_params=enable_thinking_params,
enable_memory_target=False,
enable_when_to_use=False,
enable_multiple=True,
),
ReadAllProfiles(
enable_thinking_params=enable_thinking_params,
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,
),
],
)
elif version == "halumem":
personal_summarizer = PersonalHalumemSummarizer(
tools=[
@ -248,11 +277,6 @@ class ReMe(Application):
UpdateMemoryV2(
enable_thinking_params=enable_thinking_params,
),
# RetrieveMemory(
# enable_thinking_params=enable_thinking_params,
# top_k=retrieve_top_k,
# enable_time_filter=enable_time_filter,
# ),
# 处理userprofile
ReadAllProfiles(
enable_thinking_params=enable_thinking_params,
@ -262,27 +286,19 @@ class ReMe(Application):
enable_thinking_params=enable_thinking_params,
profile_dir=self.profile_dir,
),
# AddProfile(
# enable_thinking_params=enable_thinking_params,
# profile_dir=self.profile_dir,
# ),
# DeleteProfile(
# enable_thinking_params=enable_thinking_params,
# profile_dir=self.profile_dir,
# ),
],
)
else:
raise NotImplementedError
procedural_summarizer: BaseMemoryAgent
if version in ["default", "v1", "halumem"]:
if version in ["default", "v1", "v2", "halumem"]:
procedural_summarizer = ProceduralSummarizer(tools=[])
else:
raise NotImplementedError
tool_summarizer: BaseMemoryAgent
if version in ["default", "v1", "halumem"]:
if version in ["default", "v1", "v2", "halumem"]:
tool_summarizer = ToolSummarizer(tools=[])
else:
raise NotImplementedError
@ -324,7 +340,7 @@ class ReMe(Application):
memory_agents = [personal_summarizer, procedural_summarizer, tool_summarizer]
reme_summarizer: BaseMemoryAgent
if version in ["default", "v1", "halumem"]:
if version in ["default", "v1", "v2", "halumem"]:
reme_summarizer = ReMeSummarizer(tools=[AddHistory(), DelegateTask(memory_agents=memory_agents)])
else:
raise NotImplementedError
@ -396,6 +412,29 @@ class ReMe(Application):
),
],
)
elif version == "v2":
personal_retriever = PersonalV1Retriever(
return_memory_nodes=False,
tools=[
ReadAllProfiles(
enable_thinking_params=enable_thinking_params,
enable_memory_target=False,
profile_dir=self.profile_dir,
),
RetrieveMemory(
top_k=retrieve_top_k,
enable_thinking_params=enable_thinking_params,
enable_time_filter=enable_time_filter,
enable_multiple=True,
),
ReadHistoryV2(
message_block_size=4,
vector_top_k=3,
enable_multiple=True,
enable_thinking_params=enable_thinking_params,
),
],
)
elif version == "halumem":
personal_retriever = PersonalHalumemRetriever(
tools=[
@ -408,20 +447,23 @@ class ReMe(Application):
top_k=retrieve_top_k,
enable_time_filter=enable_time_filter,
),
ReadHistory(enable_thinking_params=enable_thinking_params),
ReadHistoryV2(
message_block_size=4,
vector_top_k=3,
),
],
)
else:
raise NotImplementedError
procedural_retriever: BaseMemoryAgent
if version in ["default", "v1", "halumem"]:
if version in ["default", "v1", "v2", "halumem"]:
procedural_retriever = ProceduralRetriever(tools=[])
else:
raise NotImplementedError
tool_retriever: BaseMemoryAgent
if version in ["default", "v1", "halumem"]:
if version in ["default", "v1", "v2", "halumem"]:
tool_retriever = ToolRetriever(tools=[])
else:
raise NotImplementedError
@ -461,7 +503,7 @@ class ReMe(Application):
memory_agents = [personal_retriever, procedural_retriever, tool_retriever]
reme_retriever: BaseMemoryAgent
if version in ["default", "v1", "halumem"]:
if version in ["default", "v1", "v2", "halumem"]:
reme_retriever = ReMeRetriever(tools=[DelegateTask(memory_agents=memory_agents)])
else:
raise NotImplementedError