mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-08 22:21:15 +00:00
[dev] modify config default params
This commit is contained in:
parent
ff6b07d113
commit
b4db19ec59
6 changed files with 49 additions and 45 deletions
|
|
@ -50,8 +50,8 @@ worker:
|
|||
generation_model_top_k: 1
|
||||
retrieve_store_worker:
|
||||
class: memory.worker.read.retrieve_store_worker
|
||||
retrieve_obs_top_k: 5
|
||||
retrieve_ins_pf_top_k: 5
|
||||
retrieve_obs_top_k: 100
|
||||
retrieve_ins_pf_top_k: 100
|
||||
semantic_rank_worker:
|
||||
class: memory.worker.read.semantic_rank_worker
|
||||
fuse_rerank_worker:
|
||||
|
|
@ -60,10 +60,8 @@ worker:
|
|||
fuse_ratio_dict:
|
||||
conversation: 0.5
|
||||
observation: 1
|
||||
obs_customized: 1
|
||||
obs_customized: 1.2
|
||||
insight: 2.0
|
||||
profile: 2.0
|
||||
profile_customized: 2.0
|
||||
fuse_time_ratio: 2.0
|
||||
fuse_rerank_top_k: 10
|
||||
info_filter_worker:
|
||||
|
|
|
|||
|
|
@ -8,8 +8,4 @@ class MemoryTypeEnum(str, Enum):
|
|||
|
||||
INSIGHT = "insight"
|
||||
|
||||
PROFILE = "profile"
|
||||
|
||||
OBS_CUSTOMIZED = "obs_customized"
|
||||
|
||||
PROFILE_CUSTOMIZED = "profile_customized"
|
||||
|
|
|
|||
|
|
@ -2,9 +2,9 @@ import re
|
|||
from typing import Dict
|
||||
|
||||
from memory_scope.constants.common_constants import DATATIME_KEY_MAP, QUERY_WITH_TS, EXTRACT_TIME_DICT
|
||||
from memory_scope.constants.language_constants import DATATIME_WORD_LIST
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memory_scope.utils.datetime_handler import DatetimeHandler
|
||||
from memory_scope.utils.tool_functions import prompt_to_msg
|
||||
|
||||
|
||||
class ExtractTimeWorker(MemoryBaseWorker):
|
||||
|
|
@ -14,23 +14,21 @@ class ExtractTimeWorker(MemoryBaseWorker):
|
|||
query, query_timestamp = self.get_context(QUERY_WITH_TS)
|
||||
|
||||
# find datetime keyword
|
||||
contain_datetime = False
|
||||
for datetime_word in self.get_language_value(DATATIME_WORD_LIST):
|
||||
if datetime_word in query:
|
||||
contain_datetime = True
|
||||
break
|
||||
contain_datetime = DatetimeHandler.has_time_word(query)
|
||||
if not contain_datetime:
|
||||
self.logger.info(f"contain_datetime={contain_datetime}")
|
||||
return
|
||||
|
||||
# prepare prompt
|
||||
query_time_str = DatetimeHandler(dt=query_timestamp).string_format(self.prompt_handler.time_string_format)
|
||||
extract_time_prompt: str = self.prompt_handler.extract_time_prompt.format(query=query,
|
||||
query_time_str=query_time_str)
|
||||
self.logger.info(f"extract_time_prompt={extract_time_prompt}")
|
||||
system_prompt = self.prompt_handler.extract_time_system
|
||||
few_shot = self.prompt_handler.extract_time_few_shot.format(user_name=self.target_name)
|
||||
user_query = self.prompt_handler.extract_time_user_query.format(query=query, query_time_str=query_time_str)
|
||||
extract_time_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query)
|
||||
self.logger.info(f"extract_time_message={extract_time_message}")
|
||||
|
||||
# call sft model
|
||||
response = self.generation_model.call(prompt=extract_time_prompt, top_k=self.generation_model_top_k)
|
||||
# call llm
|
||||
response = self.generation_model.call(messages=extract_time_message, top_k=self.generation_model_top_k)
|
||||
|
||||
# if empty, return
|
||||
if not response.status or not response.message.content:
|
||||
|
|
|
|||
|
|
@ -33,9 +33,11 @@ class RetrieveStoreWorker(MemoryBaseWorker):
|
|||
|
||||
def _run(self):
|
||||
query, _ = self.get_context(QUERY_WITH_TS)
|
||||
self.submit_async_task(self.retrieve_from_observation, query=query)
|
||||
self.submit_async_task(self.retrieve_from_insight_and_profile, query=query)
|
||||
|
||||
memory_node_list: List[MemoryNode] = []
|
||||
fn_list = [self.retrieve_from_observation, self.retrieve_from_insight_and_profile]
|
||||
for result in self.async_run(fn_list=fn_list, query=query):
|
||||
for result in self.gather_async_result():
|
||||
if result:
|
||||
memory_node_list.extend(result)
|
||||
self.logger.info(f"memory_node_list.size={len(memory_node_list)}")
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from memory_scope.constants.common_constants import RETRIEVE_MEMORY_NODES, QUERY_WITH_TS, RANKED_MEMORY_NODES
|
||||
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
|
||||
|
|
@ -16,34 +15,22 @@ class SemanticRankWorker(MemoryBaseWorker):
|
|||
self.logger.warning(f"retrieve memory nodes is empty!")
|
||||
return
|
||||
|
||||
# solve content repeat, insight & profile has higher priority
|
||||
memory_node_dict: Dict[str, MemoryNode] = {}
|
||||
memory_type_selected = [MemoryTypeEnum.INSIGHT.value,
|
||||
MemoryTypeEnum.PROFILE.value,
|
||||
MemoryTypeEnum.PROFILE_CUSTOMIZED.value]
|
||||
for node in memory_node_list:
|
||||
if node.memory_type in memory_type_selected:
|
||||
return
|
||||
memory_node_dict[node.content] = node
|
||||
for node in memory_node_list:
|
||||
if node.memory_type not in memory_type_selected:
|
||||
return
|
||||
memory_node_dict[node.content] = node
|
||||
# drop repeated
|
||||
memory_node_dict: Dict[str, MemoryNode] = {n.content: n for n in memory_node_list}
|
||||
memory_node_list = list(memory_node_dict.values())
|
||||
|
||||
response = self.rank_model.call(query=query, documents=[n.content for n in memory_node_list])
|
||||
|
||||
if not response.status or not response.rank_scores:
|
||||
return
|
||||
|
||||
# set score
|
||||
rank_memory_nodes: List[MemoryNode] = []
|
||||
rank_scores = sorted(response.rank_scores.items(), key=lambda x: x[1], reverse=True)
|
||||
for idx, score in rank_scores:
|
||||
for idx, score in response.rank_scores.items():
|
||||
if idx >= len(memory_node_list):
|
||||
self.logger.warning(f"idx={idx} exceeds the maximum length of the array")
|
||||
self.logger.warning(f"idx={idx} exceeds the maximum length of rank_scores!")
|
||||
continue
|
||||
node = memory_node_list[idx]
|
||||
node.score_rank = score
|
||||
memory_node_list[idx].score_rank = score
|
||||
|
||||
memory_node_list = sorted(memory_node_list, key=lambda n: n.score_rank, reverse=True)
|
||||
for node in memory_node_list:
|
||||
self.logger.info(f"rank_stage: content={node.content} score={node.score_rank}")
|
||||
self.get_context(RANKED_MEMORY_NODES, rank_memory_nodes)
|
||||
self.get_context(RANKED_MEMORY_NODES, memory_node_list)
|
||||
|
|
|
|||
|
|
@ -2,7 +2,8 @@ import datetime
|
|||
import re
|
||||
from typing import Dict
|
||||
|
||||
from memory_scope.constants.language_constants import WEEKDAYS
|
||||
from memory_scope.constants.language_constants import WEEKDAYS, DATATIME_WORD_LIST
|
||||
from memory_scope.enumeration.language_enum import LanguageEnum
|
||||
from memory_scope.utils.global_context import G_CONTEXT
|
||||
from memory_scope.utils.logger import Logger
|
||||
|
||||
|
|
@ -96,6 +97,28 @@ class DatetimeHandler(object):
|
|||
return ""
|
||||
return getattr(cls, func_name)(extract_time_dict, meta_data)
|
||||
|
||||
@classmethod
|
||||
def has_time_word_cn(cls, query: str) -> bool:
|
||||
# find datetime keyword
|
||||
contain_datetime = False
|
||||
for datetime_word in DATATIME_WORD_LIST[LanguageEnum.CN]:
|
||||
if datetime_word in query:
|
||||
contain_datetime = True
|
||||
break
|
||||
return contain_datetime
|
||||
|
||||
@classmethod
|
||||
def has_time_word_en(cls, query: str) -> bool:
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def has_time_word(cls, query: str) -> bool:
|
||||
func_name = f"has_time_word_{G_CONTEXT.language}"
|
||||
if not hasattr(cls, func_name):
|
||||
cls.logger.warning(f"language={G_CONTEXT.language} needs to complete has_time_word func!")
|
||||
return False
|
||||
return getattr(cls, func_name)(query=query)
|
||||
|
||||
def datetime_format(self, dt_format: str = "%Y%m%d"):
|
||||
return self._dt.strftime(dt_format)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue