[dev] modify config default params

This commit is contained in:
jinli.yl 2024-07-04 00:11:43 +08:00
parent ff6b07d113
commit b4db19ec59
6 changed files with 49 additions and 45 deletions

View file

@ -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:

View file

@ -8,8 +8,4 @@ class MemoryTypeEnum(str, Enum):
INSIGHT = "insight"
PROFILE = "profile"
OBS_CUSTOMIZED = "obs_customized"
PROFILE_CUSTOMIZED = "profile_customized"

View file

@ -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:

View file

@ -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)}")

View file

@ -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)

View file

@ -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)