mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-08-28 05:25:04 +00:00
Merge branch 'main' into dev_0205
This commit is contained in:
commit
f6ca7a733a
3 changed files with 104 additions and 28 deletions
|
|
@ -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
|
||||
|
||||
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -53,7 +53,12 @@ class BaseMemoryAgent(BaseReact, metaclass=ABCMeta):
|
|||
@property
|
||||
def memory_target_type_mapping(self) -> dict[str, MemoryType]:
|
||||
"""Get the memory target type mapping from context."""
|
||||
return self.context.service_context.memory_target_type_mapping
|
||||
memory_targets = self.context.memory_targets
|
||||
memory_target_type_mapping = self.context.service_context.memory_target_type_mapping.copy()
|
||||
if memory_targets:
|
||||
return {memory_target: memory_target_type_mapping[memory_target] for memory_target in memory_targets}
|
||||
else:
|
||||
return memory_target_type_mapping
|
||||
|
||||
@property
|
||||
def meta_memory_info(self) -> str:
|
||||
|
|
|
|||
98
reme/reme.py
98
reme/reme.py
|
|
@ -26,6 +26,7 @@ from .tool.memory import (
|
|||
RetrieveMemory,
|
||||
DelegateTask,
|
||||
ReadHistory,
|
||||
ReadHistoryV2,
|
||||
ProfileHandler,
|
||||
MemoryHandler,
|
||||
AddAndRetrieveSimilarMemory,
|
||||
|
|
@ -241,6 +242,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=[
|
||||
|
|
@ -251,11 +280,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,
|
||||
|
|
@ -265,41 +289,36 @@ 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
|
||||
|
||||
memory_agents = []
|
||||
memory_targets = []
|
||||
if user_name:
|
||||
if isinstance(user_name, str):
|
||||
for message in format_messages:
|
||||
if message.role is Role.USER:
|
||||
message.name = user_name
|
||||
self._add_meta_memory(MemoryType.PERSONAL, user_name)
|
||||
memory_targets.append(user_name)
|
||||
elif isinstance(user_name, list):
|
||||
for name in user_name:
|
||||
self._add_meta_memory(MemoryType.PERSONAL, name)
|
||||
memory_targets.append(name)
|
||||
else:
|
||||
raise RuntimeError("user_name must be str or list[str]")
|
||||
memory_agents.append(personal_summarizer)
|
||||
|
|
@ -307,9 +326,11 @@ class ReMe(Application):
|
|||
if task_name:
|
||||
if isinstance(task_name, str):
|
||||
self._add_meta_memory(MemoryType.PROCEDURAL, task_name)
|
||||
memory_targets.append(task_name)
|
||||
elif isinstance(task_name, list):
|
||||
for name in task_name:
|
||||
self._add_meta_memory(MemoryType.PROCEDURAL, name)
|
||||
memory_targets.append(name)
|
||||
else:
|
||||
raise RuntimeError("task_name must be str or list[str]")
|
||||
memory_agents.append(procedural_summarizer)
|
||||
|
|
@ -317,9 +338,11 @@ class ReMe(Application):
|
|||
if tool_name:
|
||||
if isinstance(tool_name, str):
|
||||
self._add_meta_memory(MemoryType.TOOL, tool_name)
|
||||
memory_targets.append(tool_name)
|
||||
elif isinstance(tool_name, list):
|
||||
for name in tool_name:
|
||||
self._add_meta_memory(MemoryType.TOOL, name)
|
||||
memory_targets.append(name)
|
||||
else:
|
||||
raise RuntimeError("tool_name must be str or list[str]")
|
||||
memory_agents.append(tool_summarizer)
|
||||
|
|
@ -328,7 +351,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
|
||||
|
|
@ -337,6 +360,7 @@ class ReMe(Application):
|
|||
messages=format_messages,
|
||||
description=description,
|
||||
service_context=self.service_context,
|
||||
memory_targets=memory_targets,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
|
@ -400,6 +424,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=[
|
||||
|
|
@ -412,31 +459,37 @@ 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
|
||||
|
||||
memory_agents = []
|
||||
memory_targets = []
|
||||
if user_name:
|
||||
if isinstance(user_name, str):
|
||||
self._add_meta_memory(MemoryType.PERSONAL, user_name)
|
||||
memory_targets.append(user_name)
|
||||
elif isinstance(user_name, list):
|
||||
for name in user_name:
|
||||
self._add_meta_memory(MemoryType.PERSONAL, name)
|
||||
memory_targets.append(name)
|
||||
else:
|
||||
raise RuntimeError("user_name must be str or list[str]")
|
||||
memory_agents.append(personal_retriever)
|
||||
|
|
@ -444,9 +497,11 @@ class ReMe(Application):
|
|||
if task_name:
|
||||
if isinstance(task_name, str):
|
||||
self._add_meta_memory(MemoryType.PROCEDURAL, task_name)
|
||||
memory_targets.append(task_name)
|
||||
elif isinstance(task_name, list):
|
||||
for name in task_name:
|
||||
self._add_meta_memory(MemoryType.PROCEDURAL, name)
|
||||
memory_targets.append(name)
|
||||
else:
|
||||
raise RuntimeError("task_name must be str or list[str]")
|
||||
memory_agents.append(procedural_retriever)
|
||||
|
|
@ -454,9 +509,11 @@ class ReMe(Application):
|
|||
if tool_name:
|
||||
if isinstance(tool_name, str):
|
||||
self._add_meta_memory(MemoryType.TOOL, tool_name)
|
||||
memory_targets.append(tool_name)
|
||||
elif isinstance(tool_name, list):
|
||||
for name in tool_name:
|
||||
self._add_meta_memory(MemoryType.TOOL, name)
|
||||
memory_targets.append(name)
|
||||
else:
|
||||
raise RuntimeError("tool_name must be str or list[str]")
|
||||
memory_agents.append(tool_retriever)
|
||||
|
|
@ -465,7 +522,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
|
||||
|
|
@ -475,6 +532,7 @@ class ReMe(Application):
|
|||
messages=messages,
|
||||
description=description,
|
||||
service_context=self.service_context,
|
||||
memory_targets=memory_targets,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue