mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-10 22:41:06 +00:00
[dev] add info filter system prompt
This commit is contained in:
parent
64fa3badcc
commit
4333e5b13b
22 changed files with 753 additions and 99 deletions
|
|
@ -76,4 +76,13 @@ worker:
|
|||
observation: 1
|
||||
fuse_time_ratio: 2.0
|
||||
fuse_rerank_top_k: 10
|
||||
info_filter_worker:
|
||||
class: memory.worker.write.info_filter_worker
|
||||
generation_model: dashscope_generation
|
||||
info_filter_msg_max_size: 200
|
||||
generation_model_top_k: 1
|
||||
get_observation_worker:
|
||||
class: memory.worker.write.get_observation_worker
|
||||
generation_model_top_k: 1
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,10 @@
|
|||
from memory_scope.enumeration.language_enum import LanguageEnum
|
||||
|
||||
info_filter_system_prompt:
|
||||
cn: |
|
||||
任务指令:对所给{batch_size}个句子中所含有的关于用户的信息打分,分数为0,1,2或3。
|
||||
注意:其中0表示不包含用户信息,1表示句子中只包含用户假设的信息或者用户虚构的内容,2表示包含用户的一般信息,时效性信息或者需要猜测才能得到的用户信息,3表示明确含有或者可以确定推断出关于用户的重要信息,或者用户要求记录。
|
||||
按如下格式输出, 每一行输出一个打分,一定加<>,一共输出{batch_size}个分数:
|
||||
结果:
|
||||
<分数:0或1或2或3>
|
||||
|
||||
INFO_FILTER_SYSTEM_PROMPT = {
|
||||
LanguageEnum.CN: """
|
||||
|
|
@ -46,8 +46,7 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
@property
|
||||
def prompt_handler(self) -> PromptHandler:
|
||||
if self._prompt_handler is None:
|
||||
self._prompt_handler = PromptHandler()
|
||||
self._prompt_handler.add_file_prompts(self.__class__.__name__)
|
||||
self._prompt_handler = PromptHandler(__file__, **self.kwargs)
|
||||
return self._prompt_handler
|
||||
|
||||
def print_logo(self):
|
||||
|
|
|
|||
|
|
@ -83,31 +83,7 @@ TIME_MATCHED = "time_matched"
|
|||
QUERY_KEYWORDS = "query_keywords"
|
||||
|
||||
|
||||
WEEKDAYS = ["周一", "周二", "周三", "周四", "周五", "周六", "周日"]
|
||||
|
||||
DATATIME_WORD_LIST = [
|
||||
"天",
|
||||
"周",
|
||||
"月",
|
||||
"年",
|
||||
"星期",
|
||||
"点",
|
||||
"分钟",
|
||||
"小时",
|
||||
"秒",
|
||||
"上午",
|
||||
"下午",
|
||||
"早上",
|
||||
"早晨",
|
||||
"晚上",
|
||||
"中午",
|
||||
"日",
|
||||
"夜",
|
||||
"清晨",
|
||||
"傍晚",
|
||||
"凌晨",
|
||||
"岁",
|
||||
]
|
||||
|
||||
TIME_FORMAT_V1 = "{year}年{month}月{day}日{weekday}{hour}点"
|
||||
|
||||
|
|
|
|||
|
|
@ -29,3 +29,38 @@ DATATIME_WORD_LIST = {
|
|||
|
||||
]
|
||||
}
|
||||
|
||||
WEEKDAYS = {
|
||||
LanguageEnum.CN: [
|
||||
"周一",
|
||||
"周二",
|
||||
"周三",
|
||||
"周四",
|
||||
"周五",
|
||||
"周六",
|
||||
"周日"
|
||||
],
|
||||
LanguageEnum.EN: [
|
||||
|
||||
]
|
||||
}
|
||||
|
||||
NONE_WORD = {
|
||||
LanguageEnum.CN: "无",
|
||||
LanguageEnum.EN: "none"
|
||||
}
|
||||
|
||||
REPEATED_WORD = {
|
||||
LanguageEnum.CN: "重复",
|
||||
LanguageEnum.EN: "repeated"
|
||||
}
|
||||
|
||||
CONTRADICTORY_WORD = {
|
||||
LanguageEnum.CN: "矛盾",
|
||||
LanguageEnum.EN: "contradictory"
|
||||
}
|
||||
|
||||
INCLUDED_WORD = {
|
||||
LanguageEnum.CN: "被包含",
|
||||
LanguageEnum.EN: "included"
|
||||
}
|
||||
|
|
|
|||
|
|
@ -10,4 +10,5 @@ class DummyWorker(BaseWorker):
|
|||
chat_kwargs = self.get_context(CHAT_KWARGS)
|
||||
self.logger.info(f"enter workflow={workflow_name}.dummy_worker!")
|
||||
ts = int(datetime.datetime.now().timestamp())
|
||||
self.set_context(RESULT, f"test {workflow_name} kwargs={chat_kwargs} \nts={ts}")
|
||||
file_path = __file__
|
||||
self.set_context(RESULT, f"test {workflow_name} kwargs={chat_kwargs} file_path={file_path} \nts={ts}")
|
||||
|
|
|
|||
|
|
@ -83,13 +83,14 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
@property
|
||||
def prompt_handler(self) -> PromptHandler:
|
||||
if self._prompt_handler is None:
|
||||
self._prompt_handler = PromptHandler()
|
||||
self._prompt_handler.add_file_prompts(self.__class__.__name__)
|
||||
self._prompt_handler = PromptHandler(__file__, **self.kwargs)
|
||||
return self._prompt_handler
|
||||
|
||||
def __getattr__(self, key: str):
|
||||
return self.kwargs[key]
|
||||
|
||||
@staticmethod
|
||||
def get_language_prompt(prompt: dict) -> str:
|
||||
return prompt[G_CONTEXT.language]
|
||||
def get_language_value(languages: dict | list) -> str | list[str]:
|
||||
if isinstance(languages, list):
|
||||
return [x[G_CONTEXT.language] for x in languages]
|
||||
return languages[G_CONTEXT.language]
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ 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.tool_functions import time_to_formatted_str
|
||||
from memory_scope.utils.datetime_handler import DatetimeHandler
|
||||
|
||||
|
||||
class ExtractTimeWorker(MemoryBaseWorker):
|
||||
|
|
@ -15,7 +15,7 @@ class ExtractTimeWorker(MemoryBaseWorker):
|
|||
|
||||
# find datetime keyword
|
||||
contain_datetime = False
|
||||
for datetime_word in self.get_language_prompt(DATATIME_WORD_LIST):
|
||||
for datetime_word in self.get_language_value(DATATIME_WORD_LIST):
|
||||
if datetime_word in query:
|
||||
contain_datetime = True
|
||||
break
|
||||
|
|
@ -24,9 +24,7 @@ class ExtractTimeWorker(MemoryBaseWorker):
|
|||
return
|
||||
|
||||
# prepare prompt
|
||||
query_time_str = time_to_formatted_str(dt=query_timestamp,
|
||||
date_format="",
|
||||
string_format=self.prompt_handler.time_format_prompt)
|
||||
query_time_str = DatetimeHandler(dt=query_timestamp).string_format(self.prompt_handler.time_format_prompt)
|
||||
extract_time_prompt: str = self.prompt_handler.extract_time_prompt
|
||||
extract_time_prompt: str = extract_time_prompt.format(query=query, query_time_str=query_time_str)
|
||||
self.logger.info(f"extract_time_prompt={extract_time_prompt}")
|
||||
|
|
|
|||
101
memory_scope/memory/worker/write/contra_repeat_worker.py
Normal file
101
memory_scope/memory/worker/write/contra_repeat_worker.py
Normal file
|
|
@ -0,0 +1,101 @@
|
|||
from typing import List
|
||||
|
||||
from common.response_text_parser import ResponseTextParser
|
||||
from common.tool_functions import contains_keyword
|
||||
from constants.common_constants import NEW_OBS_NODES, TODAY_OBS_NODES, MSG_TIME, NEW_OBS_WITH_TIME_NODES, \
|
||||
MODIFIED_MEMORIES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class ContraRepeatWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
# 合并当前的obs和今日的obs
|
||||
new_obs_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_NODES)
|
||||
new_obs_with_time_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_WITH_TIME_NODES)
|
||||
today_obs_nodes: List[MemoryWrapNode] = self.get_context(TODAY_OBS_NODES)
|
||||
all_obs_nodes: List[MemoryWrapNode] = []
|
||||
if new_obs_nodes:
|
||||
all_obs_nodes.extend(new_obs_nodes)
|
||||
if new_obs_with_time_nodes:
|
||||
all_obs_nodes.extend(new_obs_with_time_nodes)
|
||||
if today_obs_nodes:
|
||||
all_obs_nodes.extend(today_obs_nodes)
|
||||
if not all_obs_nodes:
|
||||
self.add_run_info("all_obs_nodes is empty!")
|
||||
return
|
||||
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
all_obs_nodes = sorted(all_obs_nodes, key=lambda x: x.memory_node.metaData.get(MSG_TIME, ""), reverse=True)
|
||||
if len(all_obs_nodes) > self.config.merge_obs_max_count:
|
||||
all_obs_nodes = all_obs_nodes[:self.config.merge_obs_max_count]
|
||||
|
||||
for i, n in enumerate(all_obs_nodes):
|
||||
user_query_list.append(f"{i + 1} {n.memory_node.content}")
|
||||
|
||||
merge_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.contra_repeat_system.format(num_obs=len(user_query_list)),
|
||||
few_shot=self.prompt_config.contra_repeat_few_shot,
|
||||
user_query=self.prompt_config.contra_repeat_user_query.format(user_query="\n".join(user_query_list)))
|
||||
self.logger.info(f"merge_obs_message={merge_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text = self.gene_client.call(messages=merge_obs_message,
|
||||
model_name=self.config.merge_obs_model,
|
||||
max_token=self.config.merge_obs_max_token,
|
||||
temperature=self.config.merge_obs_temperature,
|
||||
top_k=self.config.merge_obs_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("contra repeat call llm failed!")
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat")
|
||||
if len(idx_merge_obs_list) <= 0:
|
||||
self.add_run_info("idx_merge_obs_list is empty!")
|
||||
return
|
||||
|
||||
# add merged obs
|
||||
merge_obs_nodes: List[MemoryWrapNode] = []
|
||||
for obs_content_list in idx_merge_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [6, 逃课]
|
||||
if len(obs_content_list) != 2:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, keep_flag = obs_content_list
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(all_obs_nodes):
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
if keep_flag not in ["矛盾", "被包含", "无"]:
|
||||
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
|
||||
continue
|
||||
|
||||
node: MemoryWrapNode = all_obs_nodes[idx]
|
||||
if keep_flag != "无":
|
||||
node.memory_node.status = MemoryNodeStatus.EXPIRED.value
|
||||
merge_obs_nodes.append(node)
|
||||
|
||||
# forbid keyword
|
||||
if contains_keyword(text=node.memory_node.content, keywords=self.config.forbidden_key_words):
|
||||
node.memory_node.status = MemoryNodeStatus.FORBIDDEN.value
|
||||
self.logger.info(f"after contra repeat: {node.memory_node.content} {node.memory_node.status}")
|
||||
|
||||
# save context
|
||||
self.set_context(MODIFIED_MEMORIES, merge_obs_nodes)
|
||||
29
memory_scope/memory/worker/write/es_today_obs_worker.py
Normal file
29
memory_scope/memory/worker/write/es_today_obs_worker.py
Normal file
|
|
@ -0,0 +1,29 @@
|
|||
from typing import List
|
||||
|
||||
from common.tool_functions import time_to_formatted_str
|
||||
from constants.common_constants import TODAY_OBS_NODES, DT
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsTodayObsWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
if not self.messages:
|
||||
self.logger.warning("messages is empty!")
|
||||
return
|
||||
msg_time_created = self.messages[-1].time_created
|
||||
hits = self.es_client.exact_search_v2(size=self.config.es_today_obs_top_k,
|
||||
term_filters={
|
||||
"memoryId": self.config.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": MemoryTypeEnum.OBSERVATION.value,
|
||||
f"metaData.{DT}": time_to_formatted_str(msg_time_created),
|
||||
})
|
||||
|
||||
today_obs_nodes: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(hit) for hit in hits]
|
||||
self.logger.info(f"retrieve_today_obs.size={len(today_obs_nodes)}")
|
||||
self.set_context(TODAY_OBS_NODES, today_obs_nodes)
|
||||
|
|
@ -0,0 +1,128 @@
|
|||
from datetime import datetime
|
||||
from typing import List
|
||||
|
||||
from common.response_text_parser import ResponseTextParser
|
||||
from common.tool_functions import time_to_formatted_str, get_datetime_info_dict, extract_date_parts
|
||||
from constants.common_constants import REFLECTED, DT, TIME_INFER, NEW, MSG_TIME, KEY_WORD, DATATIME_WORD_LIST, \
|
||||
NEW_OBS_WITH_TIME_NODES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from model.message import Message
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class GetObservationWithTimeWorker(MemoryBaseWorker):
|
||||
|
||||
def add_observation(self, message: Message, obs_content: str, time_infer: str, keywords: str):
|
||||
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
|
||||
dt = time_to_formatted_str(time=created_dt)
|
||||
|
||||
# 组合meta_data
|
||||
meta_data = {
|
||||
MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话
|
||||
REFLECTED: "0", # reflect标记
|
||||
DT: dt, # 当天标记
|
||||
NEW: "1", # summary-long标记
|
||||
MSG_TIME: message.time_created, # 对话时间
|
||||
TIME_INFER: time_infer, # 推断的时间
|
||||
KEY_WORD: keywords, # 关键词
|
||||
}
|
||||
|
||||
# 事件时间
|
||||
meta_data.update({f"event_{k}": str(v) for k, v in extract_date_parts(time_infer).items()})
|
||||
# 对话时间
|
||||
meta_data.update({f"msg_{k}": str(v) for k, v in get_datetime_info_dict(created_dt).items()})
|
||||
|
||||
return MemoryWrapNode.init_from_attrs(content=obs_content,
|
||||
memoryId=self.config.memory_id,
|
||||
timeCreated=message.time_created,
|
||||
scene=self.scene,
|
||||
memoryType=MemoryTypeEnum.OBSERVATION.value,
|
||||
content_modified=True, # 新增的obs需要置为true
|
||||
metaData=meta_data,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
tenantId=self.config.tenant_id)
|
||||
|
||||
def _run(self):
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
i = 1
|
||||
for msg in self.messages:
|
||||
match = False
|
||||
for time_keyword in DATATIME_WORD_LIST:
|
||||
if time_keyword in msg.content:
|
||||
match = True
|
||||
break
|
||||
if match:
|
||||
dt = time_to_formatted_str(time=msg.time_created,
|
||||
date_format="",
|
||||
string_format="{year}年{month}月{day}日{weekday}{hour}点")
|
||||
user_query_list.append(f"{i} {dt} 用户:{msg.content}")
|
||||
i += 1
|
||||
|
||||
if not user_query_list:
|
||||
self.add_run_info(f"get obs with time user_query_list={user_query_list} is empty")
|
||||
return
|
||||
|
||||
obtain_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.get_observation_with_time_system.format(num_obs=len(user_query_list)),
|
||||
few_shot=self.prompt_config.get_observation_with_time_few_shot,
|
||||
user_query=self.prompt_config.get_observation_with_time_user_query.format(
|
||||
user_query="\n".join(user_query_list)))
|
||||
self.logger.info(f"obtain_obs_message={obtain_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text: str = self.gene_client.call(messages=obtain_obs_message,
|
||||
model_name=self.config.summary_messages_model,
|
||||
max_token=self.config.summary_messages_max_token,
|
||||
temperature=self.config.summary_messages_temperature,
|
||||
top_k=self.config.summary_messages_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("summary call llm failed!", continue_run=False)
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_obs_list = ResponseTextParser(response_text).parse_v1("get_obs_time")
|
||||
if len(idx_obs_list) <= 0:
|
||||
self.add_run_info("idx_obs_list is empty!", continue_run=False)
|
||||
return
|
||||
|
||||
# gene new obs nodes
|
||||
new_obs_nodes: List[MemoryWrapNode] = []
|
||||
for obs_content_list in idx_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游]
|
||||
if len(obs_content_list) != 4:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, time_infer, obs_content, keywords = obs_content_list
|
||||
|
||||
if obs_content in ["无", "重复"]:
|
||||
continue
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
if time_infer == "无":
|
||||
time_infer = ""
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(self.messages):
|
||||
self.logger.warning(f"idx={idx} is invalid! messages.size={len(self.messages)}")
|
||||
continue
|
||||
|
||||
new_obs_nodes.append(self.add_observation(message=self.messages[idx],
|
||||
obs_content=obs_content,
|
||||
time_infer=time_infer,
|
||||
keywords=keywords))
|
||||
|
||||
# save context
|
||||
self.set_context(NEW_OBS_WITH_TIME_NODES, new_obs_nodes)
|
||||
122
memory_scope/memory/worker/write/get_observation_worker.py
Normal file
122
memory_scope/memory/worker/write/get_observation_worker.py
Normal file
|
|
@ -0,0 +1,122 @@
|
|||
from datetime import datetime
|
||||
from typing import List
|
||||
|
||||
from memory_scope.constants.common_constants import NEW_OBS_NODES, TIME_INFER
|
||||
from memory_scope.constants.language_constants import DATATIME_WORD_LIST, REPEATED_WORD, NONE_WORD
|
||||
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
|
||||
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
|
||||
from memory_scope.scheme.message import Message
|
||||
from memory_scope.utils.datetime_handler import DatetimeHandler
|
||||
from memory_scope.utils.response_text_parser import ResponseTextParser
|
||||
from memory_scope.utils.tool_functions import prompt_to_msg
|
||||
|
||||
|
||||
class GetObservationWorker(MemoryBaseWorker):
|
||||
def add_observation(self, message: Message, time_infer: str, obs_content: str, keywords: str):
|
||||
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
|
||||
dt_handler = DatetimeHandler(dt=created_dt)
|
||||
|
||||
# 组合meta_data
|
||||
meta_data = {
|
||||
MemoryTypeEnum.CONVERSATION.value: message.content,
|
||||
TIME_INFER: time_infer,
|
||||
**{f"msg_{k}": str(v) for k, v in dt_handler.dt_info_dict.items()},
|
||||
}
|
||||
|
||||
if time_infer:
|
||||
dt_infer_handler = DatetimeHandler(dt=time_infer)
|
||||
meta_data.update({f"event_{k}": str(v) for k, v in dt_infer_handler.dt_info_dict.items()})
|
||||
|
||||
node = MemoryNode(user_id=self.user_id,
|
||||
meta_data=meta_data,
|
||||
content=obs_content,
|
||||
memoryType=MemoryTypeEnum.OBSERVATION.value,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
timestamp=message.time_created,
|
||||
obs_dt=dt_handler.datetime_format(),
|
||||
obs_reflected=False,
|
||||
obs_profile_updated=False,
|
||||
keywords=keywords)
|
||||
node.gen_memory_id()
|
||||
return node
|
||||
|
||||
def build_prompt(self):
|
||||
# build prompt
|
||||
user_query_list = []
|
||||
i = 1
|
||||
for msg in self.chat_messages:
|
||||
match = False
|
||||
for time_keyword in self.get_language_value(DATATIME_WORD_LIST):
|
||||
if time_keyword in msg.content:
|
||||
match = True
|
||||
break
|
||||
if not match:
|
||||
user_query_list.append(f"{i} {self.user_id}:{msg.content}")
|
||||
i += 1
|
||||
|
||||
if not user_query_list:
|
||||
self.logger.warning(f"get obs user_query_list={user_query_list} is empty")
|
||||
return
|
||||
|
||||
system_prompt = self.prompt_handler.get_observation_system.format(num_obs=len(user_query_list),
|
||||
user_name=self.user_id)
|
||||
few_shot = self.prompt_config.get_observation_few_shot.format(self.user_id)
|
||||
user_query = self.prompt_config.get_observation_user_query.format(user_query="\n".join(user_query_list),
|
||||
user_name=self.user_id)
|
||||
|
||||
obtain_obs_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query)
|
||||
self.logger.info(f"obtain_obs_message={obtain_obs_message}")
|
||||
return obtain_obs_message
|
||||
|
||||
def _run(self):
|
||||
obtain_obs_message = self.build_prompt()
|
||||
|
||||
# call LLM
|
||||
response = self.generation_model.call(messages=obtain_obs_message, top_k=self.generation_model_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response.status or not response.message.content:
|
||||
return
|
||||
response_text = response.message.content
|
||||
|
||||
# parse text
|
||||
idx_obs_list = ResponseTextParser(response_text).parse_v1(self.__class__.__name__)
|
||||
if len(idx_obs_list) <= 0:
|
||||
self.logger.warning("idx_obs_list is empty!")
|
||||
return
|
||||
|
||||
# gene new obs nodes
|
||||
new_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [1, In June 2022, the user will travel to Hangzhou for tourism, tourism]
|
||||
if len(obs_content_list) != 4:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, time_infer, obs_content, keywords = obs_content_list
|
||||
|
||||
if obs_content in self.get_language_value([NONE_WORD, REPEATED_WORD]):
|
||||
continue
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
# index number needs to be corrected to -1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(self.messages):
|
||||
self.logger.warning(f"idx={idx} is invalid! messages.size={len(self.messages)}")
|
||||
continue
|
||||
|
||||
new_obs_nodes.append(self.add_observation(message=self.messages[idx],
|
||||
time_infer=time_infer,
|
||||
obs_content=obs_content,
|
||||
keywords=keywords))
|
||||
|
||||
# save context
|
||||
self.set_context(NEW_OBS_NODES, new_obs_nodes)
|
||||
70
memory_scope/memory/worker/write/get_observation_worker.yaml
Normal file
70
memory_scope/memory/worker/write/get_observation_worker.yaml
Normal file
|
|
@ -0,0 +1,70 @@
|
|||
get_observation_system:
|
||||
cn: |
|
||||
任务:从下面的{num_obs}句{user_name}句子中依次提取出关于{user_name}的重要信息,与相应的关键词。最多提取{num_obs}条信息。对每一句句子,只提取非常明确的信息和进行非常确定的推断,不要进行任何猜测。
|
||||
不要提取重复的信息,如果句子中的所有信息与已经提取出的信息重复了则回答“重复“,如果没有重要信息则回答“无”。注意区分,对于句子中包含{user_name}假设的信息或者{user_name}虚构的内容比如{user_name}创作的小说或剧本,不要提取信息。
|
||||
对每个句子都做一次信息提取,最后一共输出{num_obs}行信息。
|
||||
请一定要按如下格式依次输出,最后的结果一定要加<>:
|
||||
信息:<句子序号> <> <明确的重要信息或“重复”或”无“> <关键词>
|
||||
|
||||
|
||||
get_observation_few_shot:
|
||||
cn: |
|
||||
示例1:
|
||||
{user_name}句子:
|
||||
1 {user_name}:我现在处境很糟,没有工作,负债几万,怎么办
|
||||
2 {user_name}:有人说兴趣是最好的老师,也建议兴趣和职业联系起来,但我发现喜欢打篮球的人很多,但靠打篮球成职业的稀少,赚钱的更少,此外,怎么分辨兴趣和喜欢
|
||||
3 {user_name}:我现在处境很糟,没有工作,负债几万,怎么办
|
||||
4 {user_name}:我是一个刚毕业的学生,对社会,行业不了解,给我介绍一下社会系统和行业格局
|
||||
思考:从第1句可以得知{user_name}现在没有工作,负债几万,这是关于{user_name}工作与经济状况的重要信息。
|
||||
信息:<1> <> <{user_name}当前无工作且负债几万> <无工作, 负债几万>
|
||||
思考:第2句是{user_name}对他人观点的讨论和疑问,没有明确提及{user_name}个人信息。
|
||||
信息:<2> <> <无> <>
|
||||
思考:第3句含有的信息与第1句重复了。
|
||||
信息:<3> <> <重复> <>
|
||||
思考:从第4句可以得知{user_name}是一个刚毕业的学生,这是关于{user_name}身份背景状况的重要信息。其余信息重要性不足。
|
||||
信息:<4> <> <{user_name}是一名刚毕业的学生。> <刚毕业, 学生>
|
||||
|
||||
示例2:
|
||||
{user_name}句子:
|
||||
1 {user_name}:帮我写一段给同事张三女儿三岁生日的祝福语。
|
||||
2 {user_name}:能给我整理一张如何使用大模型的技巧列表吗,要求内容尽量精简。
|
||||
3 {user_name}:两个坏消息,我打羽毛球把拍子打断线了。。。然后我去我朋友家撸猫,结果我猫毛过敏,今天疯狂打喷嚏。。。
|
||||
4 {user_name}:公元1400年至1550年中国历史大事表。
|
||||
5 {user_name}:谢啦。我中午在公司附近吃,帮我推荐一家阿里巴巴徐汇滨江园区附近的餐厅吧。
|
||||
思考:从第1句可以得知张三是{user_name}的同事,这是关于{user_name}的人际关系的重要信息。其余信息重要性不足。
|
||||
信息:<1> <> <张三是{user_name}的同事。> <张三, 同事>
|
||||
思考:第2句是{user_name}提出的要求,没有明确提及{user_name}个人信息。
|
||||
信息:<2> <> <无> <>
|
||||
思考:从第3句可以得知{user_name}前天打羽毛球时把球拍打断了线,但这不是重要的信息。还可以得知{user_name}对猫毛过敏,这是关于{user_name}的健康的重要信息。
|
||||
信息:<3> <> <{user_name}对猫毛过敏。> <猫毛, 过敏>
|
||||
思考:从第4句是{user_name}提出的要求,没有明确提及{user_name}个人信息。
|
||||
信息:<4> <> <无> <>
|
||||
思考:从第5句可以得知{user_name}在阿里巴巴徐汇滨江园区工作,这是关于{user_name}的工作地点的重要信息。
|
||||
信息:<5> <> <{user_name}在阿里巴巴徐汇滨江园区工作。> <阿里巴巴, 徐汇滨江园区, 工作>
|
||||
|
||||
示例3:
|
||||
{user_name}句子:
|
||||
1 {user_name}:我想买辆新能源汽车,有什么推荐吗?
|
||||
2 {user_name}:我在上海,想买辆新能源汽车,有什么推荐吗?
|
||||
3 {user_name}:案外人异议审查期间,人民法院不得对执行标的进行处分,不就是中止执行的意思吗?
|
||||
4 {user_name}:请写两句藏头诗分别以“胜”和“利”开头。
|
||||
5 {user_name}:我花5000元买了100股海天味业。
|
||||
6 {user_name}:李增杰:这个是星座蛙设,但是我是处女座的,我妈感觉因为我的不正常,我妈不让我看了\n雌猴摸了摸李增杰的头,这样啊\n雌猴打开了哔哩哔哩看了看\n雌猴:要不换个设吧,我听你未来的你说,有一个叫难忘的朱古力232这个人,他弄的设是Windows设\n这是剧本1,剧本2未完待续
|
||||
思考:从第1句可以得知{user_name}寻求购买新能源汽车的建议或推荐,这是这是关于{user_name}的大宗消费的重要的信息。
|
||||
信息:<1> <> <{user_name}寻求购买新能源汽车的建议或推荐。> <购买, 新能源汽车>
|
||||
思考:从第2句可以得知{user_name}当前所在城市为上海,这是关于{user_name}的生活地区的重要信息。其余信息与第1句重复了。
|
||||
信息:<2> <> <{user_name}所在的城市是上海。> <上海>
|
||||
思考:第3句是{user_name}对某个观点的讨论和疑问,没有明确提及{user_name}个人信息。
|
||||
信息:<3> <> <无> <>
|
||||
思考:第4句是{user_name}提出的要求,没有明确提及{user_name}个人信息。
|
||||
信息:<4> <> <无> <>
|
||||
思考:从第5句可以得知{user_name}购买了海天味业股票,购买数量为100股,购买金额为5000元,这是关于{user_name}的投资决策的重要信息。
|
||||
信息:<5> <> <{user_name}购买了海天味业股票,购买数量为100股,购买金额为5000元。> <海天味业, 股票>
|
||||
思考:第6句是{user_name}创作的剧本内容,无法提取{user_name}个人信息。
|
||||
信息:<6> <> <无> <>
|
||||
|
||||
|
||||
get_observation_user_query:
|
||||
cn: |
|
||||
{user_name}句子:
|
||||
{user_query}
|
||||
58
memory_scope/memory/worker/write/info_filter_worker.py
Normal file
58
memory_scope/memory/worker/write/info_filter_worker.py
Normal file
|
|
@ -0,0 +1,58 @@
|
|||
from typing import List
|
||||
|
||||
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memory_scope.scheme.message import Message
|
||||
from memory_scope.utils.response_text_parser import ResponseTextParser
|
||||
from memory_scope.utils.tool_functions import prompt_to_msg
|
||||
|
||||
|
||||
class InfoFilterWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
# filter user msg
|
||||
info_messages: List[Message] = []
|
||||
for msg in self.chat_messages:
|
||||
# TODO: add memory for assistant
|
||||
if msg.role != MessageRoleEnum.USER.value:
|
||||
continue
|
||||
if len(msg.content) >= self.info_filter_msg_max_size:
|
||||
begin_size = int(self.info_filter_msg_max_size * 0.75 + 0.5)
|
||||
end_size = int(self.info_filter_msg_max_size * 0.25 + 0.5)
|
||||
msg.content = msg.content[: begin_size] + msg.content[-end_size:]
|
||||
info_messages.append(msg)
|
||||
|
||||
# gene prompt
|
||||
user_query = "\n".join([f"{i + 1} {self.user_id}:{msg.content}" for i, msg in enumerate(info_messages)])
|
||||
system_prompt = self.prompt_handler.info_filter_system.format(batch_size=len(info_messages),
|
||||
user_name=self.user_id)
|
||||
few_shot = self.prompt_handler.info_filter_few_shot.format(user_name=self.user_id)
|
||||
user_query = self.prompt_handler.info_filter_user_query.format(user_query=user_query)
|
||||
info_filter_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query)
|
||||
self.logger.info(f"info_filter_message={info_filter_message}")
|
||||
|
||||
# call llm
|
||||
response = self.generation_model.call(messages=info_filter_message, top_k=self.generation_model_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response.status or not response.message.content:
|
||||
return
|
||||
response_text = response.message.content
|
||||
|
||||
# parse text
|
||||
info_score_list = ResponseTextParser(response_text).parse_v1(self.__class__.__name__)
|
||||
if len(info_score_list) != len(info_messages):
|
||||
self.logger.warning(f"score_size != messages_size, {len(info_score_list)} vs {len(info_messages)}")
|
||||
return
|
||||
|
||||
# filter messages
|
||||
filtered_messages: List[Message] = []
|
||||
for msg, info_score in zip(info_messages, info_score_list):
|
||||
if not info_score:
|
||||
continue
|
||||
|
||||
score = info_score[0]
|
||||
if score in ("3",):
|
||||
msg.meta_data["info_score"] = score
|
||||
filtered_messages.append(msg)
|
||||
self.chat_messages = filtered_messages
|
||||
63
memory_scope/memory/worker/write/info_filter_worker.yaml
Normal file
63
memory_scope/memory/worker/write/info_filter_worker.yaml
Normal file
|
|
@ -0,0 +1,63 @@
|
|||
info_filter_system:
|
||||
cn: |
|
||||
任务指令:对所给{batch_size}个句子中所含有的关于{user_name}的信息打分,分数为0,1,2或3。
|
||||
注意:其中0表示不包含{user_name}信息,1表示句子中只包含{user_name}假设的信息或者{user_name}虚构的内容,2表示包含{user_name}的一般信息,时效性信息或者需要猜测才能得到的{user_name}信息,3表示明确含有或者可以确定推断出关于{user_name}的重要信息,或者{user_name}要求记录。
|
||||
按如下格式输出, 每一行输出一个打分,一定加<>,一共输出{batch_size}个分数:
|
||||
结果:
|
||||
<分数:0或1或2或3>
|
||||
|
||||
info_filter_few_shot:
|
||||
cn: |
|
||||
示例1
|
||||
句子:
|
||||
1 {user_name}:帮我写一段给同事张三女儿三岁生日的祝福语。
|
||||
2 {user_name}:公元1400年至1550年中国历史大事表。
|
||||
3 {user_name}:你吃午饭了吗?
|
||||
4 {user_name}:我今天心情不好,可以安慰我一下吗?
|
||||
5 {user_name}:能给我整理一张如何使用大模型的技巧列表吗,要求内容尽量精简。
|
||||
6 {user_name}:记一下,明天下午3点提醒我去拿一下文件。
|
||||
结果:
|
||||
<3>
|
||||
<0>
|
||||
<0>
|
||||
<2>
|
||||
<2>
|
||||
<3>
|
||||
|
||||
示例2
|
||||
句子:
|
||||
1 {user_name}:我刚刚入职了阿里巴巴。
|
||||
2 {user_name}:露天睡觉蚊子多,咋搞。
|
||||
3 {user_name}:创造力和外倾性有关?
|
||||
4 {user_name}:一个区县的所有的事业人员的档案审核、修改和规范,应该是县委组织部下属的干部档案中心负责还是县人社局负责?
|
||||
5 {user_name}:假如我要和一个女人准备要孩子,我作为男人,怎么保护女人和孩子以及怎么备孕确保精子质量高对后代好
|
||||
6 {user_name}:我和你一起出去玩,你会感觉开心吗?
|
||||
7 {user_name}:林浅,一位对未来充满好奇的年轻女孩,偶然间发现了这家能寄信给未来的邮局。出于对逝去祖父的怀念,她决定写下一封信,寄给五年后的自己,希望能收到祖父生前未说完的故事。五年期限将至,当她几乎忘记这段往事时,一封泛黄的回信悄然降临,不仅带来了祖父未完的冒险故事,还藏着一段关于勇气、爱与自我发现的深刻启示。续写成3000字小说。
|
||||
结果:
|
||||
<3>
|
||||
<2>
|
||||
<0>
|
||||
<0>
|
||||
<3>
|
||||
<1>
|
||||
<1>
|
||||
|
||||
示例3
|
||||
句子:
|
||||
1 {user_name}:你的妈妈患有焦虑症,怎么安慰和开导她?
|
||||
2 {user_name}:肾脏严重亏空
|
||||
3 {user_name}:我很喜欢打篮球,所以我身体很好
|
||||
4 {user_name}:篮球明星有哪些?
|
||||
5 {user_name}:李增杰:这个是星座蛙设,但是我是处女座的,我妈感觉因为我的不正常,我妈不让我看了\n雌猴摸了摸李增杰的头,这样啊\n雌猴打开了哔哩哔哩看了看\n雌猴:要不换个设吧,我听你未来的你说,有一个叫难忘的朱古力232这个人,他弄的设是Windows设\n这是剧本1,剧本2未完待续
|
||||
结果:
|
||||
<1>
|
||||
<1>
|
||||
<3>
|
||||
<0>
|
||||
<1>
|
||||
|
||||
info_filter_user_query:
|
||||
cn: |
|
||||
句子:
|
||||
{user_query}
|
||||
结果:
|
||||
|
|
@ -11,6 +11,8 @@ class MemoryNode(BaseModel):
|
|||
|
||||
user_id: str = Field("", description="unique memory id for user")
|
||||
|
||||
meta_data: Dict[str, str] = Field({}, description="other data infos")
|
||||
|
||||
content: str = Field("", description="memory content")
|
||||
|
||||
score_similar: float = Field(0, description="es similar score")
|
||||
|
|
@ -21,14 +23,21 @@ class MemoryNode(BaseModel):
|
|||
|
||||
memory_type: str = Field("", description="conversation/observation/insight...")
|
||||
|
||||
meta_data: Dict[str, str] = Field({}, description="other data infos")
|
||||
|
||||
status: str = Field("active", description="active or expired")
|
||||
|
||||
vector: List[float] = Field([], description="content embedding result, return empty")
|
||||
|
||||
timestamp: int = Field(int(datetime.datetime.now().timestamp()), description="timestamp of the memory node")
|
||||
|
||||
obs_dt: str = Field("", description="dt of the observation")
|
||||
|
||||
obs_reflected: bool = Field(False, description="if the observation is reflected")
|
||||
|
||||
obs_profile_updated: bool = Field(False, description="if the observation has updated user profile")
|
||||
|
||||
keyword: str = Field("", description="keywords of the content")
|
||||
|
||||
|
||||
@property
|
||||
def node_keys(self):
|
||||
return list(self.model_json_schema()["properties"].keys())
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ class DummyVectorStore(BaseVectorStore):
|
|||
|
||||
def __init__(self, embedding_model: BaseModel, **kwargs):
|
||||
self.embedding_model: BaseModel = embedding_model
|
||||
self.kwargs = kwargs
|
||||
|
||||
def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
|
||||
pass
|
||||
|
|
|
|||
78
memory_scope/utils/datetime_handler.py
Normal file
78
memory_scope/utils/datetime_handler.py
Normal file
|
|
@ -0,0 +1,78 @@
|
|||
import datetime
|
||||
import re
|
||||
|
||||
from memory_scope.constants.language_constants import WEEKDAYS
|
||||
from memory_scope.utils.global_context import G_CONTEXT
|
||||
from memory_scope.utils.logger import Logger
|
||||
|
||||
|
||||
class DatetimeHandler(object):
|
||||
|
||||
def __init__(self, dt: datetime.datetime | str | int | float = None):
|
||||
if isinstance(dt, str | int | float):
|
||||
if isinstance(dt, str):
|
||||
dt = float(dt)
|
||||
self._dt: datetime.datetime = datetime.datetime.fromtimestamp(dt)
|
||||
elif isinstance(dt, datetime.datetime):
|
||||
self._dt: datetime.datetime = dt
|
||||
else:
|
||||
self._dt: datetime.datetime = datetime.datetime.now()
|
||||
|
||||
self._dt_info_dict: dict | None = None
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
def _parse_dt_info(self):
|
||||
return {
|
||||
"year": self._dt.year,
|
||||
"month": self._dt.month,
|
||||
"day": self._dt.day,
|
||||
"hour": self._dt.hour,
|
||||
"minute": self._dt.minute,
|
||||
"second": self._dt.second,
|
||||
"week": self._dt.isocalendar().week,
|
||||
"weekday": WEEKDAYS[G_CONTEXT.language][self._dt.isocalendar().weekday - 1],
|
||||
}
|
||||
|
||||
@property
|
||||
def dt_info_dict(self):
|
||||
if self._dt_info_dict is None:
|
||||
self._dt_info_dict = self._parse_dt_info()
|
||||
return self._dt_info_dict
|
||||
|
||||
@staticmethod
|
||||
def extract_date_parts_cn(input_string: str):
|
||||
# Extending our pattern to handle every/每 as a possible value.
|
||||
patterns = {
|
||||
'year': r'(\d+|每)年',
|
||||
'month': r'(\d+|每)月',
|
||||
'day': r'(\d+|每)日',
|
||||
'weekday': r'周([一二三四五六日])',
|
||||
'hour': r'(\d+)点'
|
||||
}
|
||||
weekday_dict = {"一": 1, "二": 2, "三": 3, "四": 4, "五": 5, "六": 6, "日": 7}
|
||||
extracted_data = {}
|
||||
|
||||
# Search for patterns in the input string and populate the dictionary
|
||||
for key, pattern in patterns.items():
|
||||
match = re.search(pattern, input_string)
|
||||
if match: # If there is a match, include it in the output dictionary
|
||||
if match.group(1) == "每":
|
||||
extracted_data[key] = -1
|
||||
elif match.group(1) in weekday_dict.keys():
|
||||
extracted_data[key] = weekday_dict[match.group(1)]
|
||||
else:
|
||||
extracted_data[key] = int(match.group(1))
|
||||
return extracted_data
|
||||
|
||||
def extract_date_parts(self):
|
||||
func_name = f"extract_date_parts_{G_CONTEXT.language}"
|
||||
if not hasattr(self, func_name):
|
||||
self.logger.warning(f"language={G_CONTEXT.language} needs to complete extract_date_parts function!")
|
||||
return {}
|
||||
return getattr(self, func_name)()
|
||||
|
||||
def datetime_format(self, dt_format: str = "%Y%m%d"):
|
||||
return self._dt.strftime(dt_format)
|
||||
|
||||
def string_format(self, string_format: str):
|
||||
return string_format.format(**self.dt_info_dict)
|
||||
|
|
@ -5,30 +5,37 @@ from typing import Dict
|
|||
import yaml
|
||||
|
||||
from memory_scope.utils.global_context import G_CONTEXT
|
||||
from memory_scope.utils.tool_functions import camelcase_to_underscore
|
||||
|
||||
|
||||
class PromptHandler(object):
|
||||
|
||||
def __init__(self, default_prompt_dir: str = "config/prompts"):
|
||||
self._default_prompt_dir: str = default_prompt_dir
|
||||
def __init__(self, class_path: str, prompt_file: str = "", prompt_dict: dict = None, **kwargs):
|
||||
self._class_path: str = class_path
|
||||
self._prompt_dict: Dict[str, str] = {}
|
||||
|
||||
def add_file_prompts(self, name: str, to_underscore: bool = True):
|
||||
if to_underscore:
|
||||
name: str = camelcase_to_underscore(name)
|
||||
file_path = self._class_path.strip(".py")
|
||||
self.add_prompt_file(file_path)
|
||||
|
||||
class_path = os.path.join(self._default_prompt_dir, name)
|
||||
if os.path.exists(f"{class_path}.yaml"):
|
||||
with open(f"{class_path}.yaml") as f:
|
||||
prompt_language_dict = yaml.load(f, yaml.FullLoader)
|
||||
elif os.path.exists(f"{class_path}.json"):
|
||||
with open(f"{class_path}.json") as f:
|
||||
prompt_language_dict = json.load(f)
|
||||
if prompt_file:
|
||||
self.add_prompt_file(prompt_file)
|
||||
|
||||
if prompt_dict:
|
||||
self.add_prompt_dict(prompt_dict)
|
||||
|
||||
def add_prompt_file(self, file_path: str):
|
||||
if os.path.exists(f"{file_path}.yaml"):
|
||||
with open(f"{file_path}.yaml") as f:
|
||||
prompt_dict = yaml.load(f, yaml.FullLoader)
|
||||
elif os.path.exists(f"{file_path}.json"):
|
||||
with open(f"{file_path}.json") as f:
|
||||
prompt_dict = json.load(f)
|
||||
else:
|
||||
raise RuntimeError(f"{class_path}.yaml/json is not exists!")
|
||||
raise RuntimeError(f"{file_path}.yaml/json is not exists!")
|
||||
|
||||
for key, language_dict in prompt_language_dict.items():
|
||||
self.add_prompt_dict(prompt_dict)
|
||||
|
||||
def add_prompt_dict(self, prompt_dict: dict):
|
||||
for key, language_dict in prompt_dict.items():
|
||||
prompts = language_dict.get(G_CONTEXT.language)
|
||||
if not prompts:
|
||||
raise RuntimeError(f"{key}.prompt.{G_CONTEXT.language} is empty!")
|
||||
|
|
|
|||
|
|
@ -9,8 +9,9 @@ from importlib import import_module
|
|||
import pyfiglet
|
||||
from termcolor import colored, COLORS
|
||||
|
||||
from memory_scope.constants.common_constants import WEEKDAYS
|
||||
from memory_scope.constants.language_constants import WEEKDAYS
|
||||
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memory_scope.utils.global_context import G_CONTEXT
|
||||
|
||||
|
||||
def underscore_to_camelcase(name: str, is_first_title: bool = True):
|
||||
|
|
@ -25,11 +26,15 @@ def camelcase_to_underscore(name: str):
|
|||
return re.sub(r'(?<!^)(?=[A-Z])', '_', name).lower()
|
||||
|
||||
|
||||
def init_instance_by_config(config: dict, default_class_path: str = "memory_scope", suffix_name: str = "", **kwargs):
|
||||
def init_instance_by_config(config: dict,
|
||||
default_class_path: str = "memory_scope",
|
||||
suffix_name: str = "",
|
||||
**kwargs):
|
||||
config_copy = deepcopy(config)
|
||||
origin_class_path: str = config_copy.pop("class")
|
||||
if not origin_class_path:
|
||||
raise RuntimeError("empty class path!")
|
||||
user_defined: bool = config_copy.pop("user_defined", False)
|
||||
|
||||
class_name_split = origin_class_path.split(".")
|
||||
class_name: str = class_name_split[-1]
|
||||
|
|
@ -38,7 +43,7 @@ def init_instance_by_config(config: dict, default_class_path: str = "memory_scop
|
|||
class_name_split[-1] = class_name
|
||||
|
||||
class_paths = []
|
||||
if default_class_path and not origin_class_path.startswith(default_class_path):
|
||||
if not user_defined and default_class_path and not origin_class_path.startswith(default_class_path):
|
||||
class_paths.append(default_class_path)
|
||||
class_paths.extend(class_name_split)
|
||||
module = import_module(".".join(class_paths))
|
||||
|
|
@ -48,12 +53,6 @@ def init_instance_by_config(config: dict, default_class_path: str = "memory_scop
|
|||
return getattr(module, cls_name)(**config_copy)
|
||||
|
||||
|
||||
def complete_config_name(config_name: str, suffix: str = ".json"):
|
||||
if not config_name.endswith(suffix):
|
||||
config_name += suffix
|
||||
return config_name
|
||||
|
||||
|
||||
def prompt_to_msg(system_prompt: str, few_shot: str, user_query: str):
|
||||
return [
|
||||
{
|
||||
|
|
@ -69,41 +68,6 @@ def prompt_to_msg(system_prompt: str, few_shot: str, user_query: str):
|
|||
]
|
||||
|
||||
|
||||
def get_datetime_info_dict(parse_dt: datetime):
|
||||
return {
|
||||
"year": parse_dt.year,
|
||||
"month": parse_dt.month,
|
||||
"day": parse_dt.day,
|
||||
"hour": parse_dt.hour,
|
||||
"minute": parse_dt.minute,
|
||||
"second": parse_dt.second,
|
||||
"week": parse_dt.isocalendar().week,
|
||||
"weekday": WEEKDAYS[parse_dt.isocalendar().weekday - 1],
|
||||
}
|
||||
|
||||
|
||||
def time_to_formatted_str(dt: datetime | str | int | float = None,
|
||||
date_format: str = "%Y%m%d", # e.g. %Y%m%d -> "20240528", add %H:%M:%S
|
||||
string_format: str = "") -> str:
|
||||
|
||||
if isinstance(dt, str | int | float):
|
||||
if isinstance(dt, str):
|
||||
dt = float(dt)
|
||||
current_dt = datetime.fromtimestamp(dt)
|
||||
elif isinstance(dt, datetime):
|
||||
current_dt = dt
|
||||
else:
|
||||
current_dt = datetime.now()
|
||||
|
||||
return_str = ""
|
||||
if date_format:
|
||||
return_str = current_dt.strftime(date_format)
|
||||
elif string_format:
|
||||
return_str = string_format.format(**get_datetime_info_dict(current_dt))
|
||||
|
||||
return return_str
|
||||
|
||||
|
||||
def char_logo(words: str, seed: int = time.time_ns(), color=None):
|
||||
font = pyfiglet.Figlet()
|
||||
rendered_text = font.renderText(words)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue