ReMe/old/retrieve/extract_time_worker.py
2024-06-20 22:18:57 +08:00

71 lines
3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import re
from common.tool_functions import time_to_formatted_str
from constants.common_constants import DATATIME_WORD_LIST, DATATIME_KEY_MAP
from constants.common_constants import EXTRACT_TIME_DICT
from worker.memory_base_worker import MemoryBaseWorker
class ExtractTimeWorker(MemoryBaseWorker):
def __init__(self, parse_time_model, parse_time_max_token, parse_time_temperature, parse_time_top_k, *args, **kwargs):
super(ExtractTimeWorker, self).__init__(*args, **kwargs)
self.parse_time_model = parse_time_model
self.parse_time_max_token = parse_time_max_token
self.parse_time_temperature = parse_time_temperature
self.parse_time_top_k = parse_time_top_k
@staticmethod
def get_parse_time_prompt(query: str, query_time_str: str):
return f"""
任务指令:从语句与语句发生的时间,推断并提取语句内容中指向的时间段。回答尽可能完整的时间段。
语句:{query}
时间:{query_time_str}
回答:
""".strip()
def _run(self):
# save to context
extract_time_dict = {}
self.set_context(EXTRACT_TIME_DICT, extract_time_dict)
# get query & time_created_dt
query = self.messages[-1].content
time_created = self.messages[-1].time_created
# find datetime keyword
contain_datetime = False
for datetime_word in DATATIME_WORD_LIST:
if datetime_word in query:
contain_datetime = True
break
if not contain_datetime:
self.logger.info(f"contain_datetime={contain_datetime}")
return
# prepare prompt
time_format = "{year}{month}{day}日,{year}年第{week}周,{weekday}{hour}{minute}{second}秒。"
query_time_str = time_to_formatted_str(time=time_created,
date_format="",
string_format=time_format)
extract_time_prompt = self.get_parse_time_prompt(query=query, query_time_str=query_time_str)
self.logger.info(f"extract_time_prompt={extract_time_prompt}")
# call sft model
response_text = self.gene_client.call(prompt=extract_time_prompt,
model_name=self.parse_time_model,
max_token=self.parse_time_max_token,
temperature=self.parse_time_temperature,
top_k=self.parse_time_top_k)
# if empty, return
if not response_text:
return
# re-match time info to dict
pattern = r'-\s*(\S+)(\d+)'
matches = re.findall(pattern, response_text)
for key, value in matches:
if key in DATATIME_KEY_MAP.keys():
extract_time_dict[DATATIME_KEY_MAP[key]] = value
self.logger.info(f"response_text={response_text} filters={extract_time_dict}")