From b4db19ec59d861ee030991f624b5e514c7eae3a1 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 4 Jul 2024 00:11:43 +0800 Subject: [PATCH] [dev] modify config default params --- config/demo_config.yaml | 8 ++--- memory_scope/enumeration/memory_type_enum.py | 4 --- .../memory/worker/read/extract_time_worker.py | 20 ++++++------ .../worker/read/retrieve_store_worker.py | 6 ++-- .../worker/read/semantic_rank_worker.py | 31 ++++++------------- memory_scope/utils/datetime_handler.py | 25 ++++++++++++++- 6 files changed, 49 insertions(+), 45 deletions(-) diff --git a/config/demo_config.yaml b/config/demo_config.yaml index 3cca4457..5b6715bf 100644 --- a/config/demo_config.yaml +++ b/config/demo_config.yaml @@ -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: diff --git a/memory_scope/enumeration/memory_type_enum.py b/memory_scope/enumeration/memory_type_enum.py index e4e72b7e..42a4e413 100644 --- a/memory_scope/enumeration/memory_type_enum.py +++ b/memory_scope/enumeration/memory_type_enum.py @@ -8,8 +8,4 @@ class MemoryTypeEnum(str, Enum): INSIGHT = "insight" - PROFILE = "profile" - OBS_CUSTOMIZED = "obs_customized" - - PROFILE_CUSTOMIZED = "profile_customized" diff --git a/memory_scope/memory/worker/read/extract_time_worker.py b/memory_scope/memory/worker/read/extract_time_worker.py index 4416eb93..2026da8b 100644 --- a/memory_scope/memory/worker/read/extract_time_worker.py +++ b/memory_scope/memory/worker/read/extract_time_worker.py @@ -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: diff --git a/memory_scope/memory/worker/read/retrieve_store_worker.py b/memory_scope/memory/worker/read/retrieve_store_worker.py index d685d826..5cdc3ab6 100644 --- a/memory_scope/memory/worker/read/retrieve_store_worker.py +++ b/memory_scope/memory/worker/read/retrieve_store_worker.py @@ -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)}") diff --git a/memory_scope/memory/worker/read/semantic_rank_worker.py b/memory_scope/memory/worker/read/semantic_rank_worker.py index 161b6e2f..5efa1571 100644 --- a/memory_scope/memory/worker/read/semantic_rank_worker.py +++ b/memory_scope/memory/worker/read/semantic_rank_worker.py @@ -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) diff --git a/memory_scope/utils/datetime_handler.py b/memory_scope/utils/datetime_handler.py index b0384644..2f1949a3 100644 --- a/memory_scope/utils/datetime_handler.py +++ b/memory_scope/utils/datetime_handler.py @@ -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)