[dev] reformat llm params

This commit is contained in:
jinli.yl 2024-07-01 23:07:36 +08:00
parent b151da319e
commit c2406e2eba
8 changed files with 75 additions and 43 deletions

View file

@ -28,7 +28,8 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
self._vector_store: BaseVectorStore | None = None
self._monitor: BaseMonitor | None = None
self._user_id: str | None = None
self._user_name: str | None = None
self._target_name: str | None = None
self._prompt_handler: PromptHandler | None = None
@property
@ -74,11 +75,20 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
return self._monitor
@property
def user_id(self) -> str:
if self._user_id is None:
def user_name(self) -> str:
# FIXME: complex situations require adjustment
if self._user_name is None:
message = [x for x in self.messages if x.role == MessageRoleEnum.ASSISTANT.value][-1]
self._user_name = message.role_name
return self._user_name
@property
def target_name(self) -> str:
# FIXME: complex situations require adjustment
if self._target_name is None:
message = [x for x in self.messages if x.role == MessageRoleEnum.USER.value][-1]
self._user_id = message.role_name
return self._user_id
self._target_name = message.role_name
return self._target_name
@property
def prompt_handler(self) -> PromptHandler:

View file

@ -30,11 +30,7 @@ class ExtractTimeWorker(MemoryBaseWorker):
self.logger.info(f"extract_time_prompt={extract_time_prompt}")
# call sft model
response = self.generation_model.call(prompt=extract_time_prompt,
model_name=self.extra_time_model,
max_token=self.extra_time_max_token,
temperature=self.extra_time_temperature,
top_k=self.extra_time_top_k)
response = self.generation_model.call(prompt=extract_time_prompt, top_k=self.extra_time_top_k)
# if empty, return
if not response.status or not response.message.content:
@ -47,6 +43,5 @@ class ExtractTimeWorker(MemoryBaseWorker):
for key, value in matches:
if key in DATATIME_KEY_MAP.keys():
extract_time_dict[DATATIME_KEY_MAP[key]] = value
self.set_context(EXTRACT_TIME_DICT, extract_time_dict)
self.logger.info(f"response_text={response_text} filters={extract_time_dict}")
self.set_context(EXTRACT_TIME_DICT, extract_time_dict)

View file

@ -1,8 +1,10 @@
from typing import Dict, List
from memory_scope.constants.common_constants import EXTRACT_TIME_DICT, RANKED_MEMORY_NODES, RESULT
from memory_scope.constants.language_constants import COLON_WORD
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.utils.datetime_handler import DatetimeHandler
class FuseRerankWorker(MemoryBaseWorker):
@ -107,10 +109,10 @@ class FuseRerankWorker(MemoryBaseWorker):
content = node.content
if f_event or f_msg:
time_infer = self.format_time_infer(time_infer="",
extract_time_dict=extract_time_dict,
meta_data=node.memory_node.metaData)
content = f"{time_infer}: {content}"
time_infer = DatetimeHandler.format_time_by_extract_time(extract_time_dict=extract_time_dict,
meta_data=node.memory_node.metaData)
if time_infer:
content = f"{time_infer}{self.get_language_value(COLON_WORD)}{content}"
memories.append(content)
self.set_context(RESULT, "\n".join(memories))

View file

@ -11,7 +11,8 @@ class RetrieveStoreWorker(MemoryBaseWorker):
async def retrieve_from_observation(self, query: str) -> List[MemoryNode]:
filter_dict = {
"user_id": self.user_id,
"user_name": self.user_name,
"target_name": self.target_name,
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
}
@ -21,7 +22,8 @@ class RetrieveStoreWorker(MemoryBaseWorker):
async def retrieve_from_insight_and_profile(self, query: str) -> List[MemoryNode]:
filter_dict = {
"user_id": self.user_id,
"user_name": self.user_name,
"target_name": self.target_name,
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": [MemoryTypeEnum.INSIGHT.value, MemoryTypeEnum.PROFILE.value],
}

View file

@ -1,4 +1,3 @@
from datetime import datetime
from typing import List
from memory_scope.constants.common_constants import NEW_OBS_NODES, TIME_INFER
@ -15,21 +14,21 @@ 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)
dt_handler = DatetimeHandler(dt=message.time_created)
# 组合meta_data
# buidl 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()},
**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()})
dt_info_dict = DatetimeHandler.extract_date_parts(input_string=time_infer)
meta_data.update({f"event_{k}": str(v) for k, v in dt_info_dict.items()})
node = MemoryNode(user_id=self.user_id,
node = MemoryNode(user_name=self.user_name,
target_name=self.target_name,
meta_data=meta_data,
content=obs_content,
memoryType=MemoryTypeEnum.OBSERVATION.value,
@ -38,7 +37,7 @@ class GetObservationWorker(MemoryBaseWorker):
obs_dt=dt_handler.datetime_format(),
obs_reflected=False,
obs_profile_updated=False,
keywords=keywords)
obs_keyword=keywords)
node.gen_memory_id()
return node

View file

@ -7,11 +7,13 @@ from memory_scope.utils.tool_functions import md5_hash
class MemoryNode(BaseModel):
memory_id: str = Field("", description="unique id for memory item")
memory_id: str = Field("", description="unique id for memory")
user_id: str = Field("", description="unique memory id for user")
user_name: str = Field("", description="the user who owns the memory")
meta_data: Dict[str, str] = Field({}, description="other data infos")
target_name: str = Field("", description="target name described by the memory")
meta_data: Dict[str, str] = Field({}, description="meta data infos")
content: str = Field("", description="memory content")
@ -35,8 +37,7 @@ class MemoryNode(BaseModel):
obs_profile_updated: bool = Field(False, description="if the observation has updated user profile")
keyword: str = Field("", description="keywords of the content")
obs_keyword: str = Field("", description="keywords of the content")
@property
def node_keys(self):
@ -46,4 +47,4 @@ class MemoryNode(BaseModel):
return self.model_dump().get(key)
def gen_memory_id(self):
self.memory_id = f"{self.user_id}_{self.timestamp}_{md5_hash(self.content)[:8]}"
self.memory_id = f"{self.user_name}_{self.target_name}_{self.timestamp}_{md5_hash(self.content)[:8]}"

View file

@ -92,9 +92,7 @@ class LlamaIndexElasticSearchStore(BaseVectorStore):
filter_dict = {}
es_filter = _to_elasticsearch_filter(filter_dict)
retriever = self.index.as_retriever(
vector_store_kwargs={"es_filter": es_filter},
similarity_top_k=top_k)
retriever = self.index.as_retriever(vector_store_kwargs={"es_filter": es_filter}, similarity_top_k=top_k)
text_nodes = retriever.retrieve(query)
return [self._text_node_2_memory_node(n) for n in text_nodes]

View file

@ -1,5 +1,6 @@
import datetime
import re
from typing import Dict
from memory_scope.constants.language_constants import WEEKDAYS
from memory_scope.utils.global_context import G_CONTEXT
@ -7,6 +8,7 @@ from memory_scope.utils.logger import Logger
class DatetimeHandler(object):
logger = Logger.get_logger()
def __init__(self, dt: datetime.datetime | str | int | float = None):
if isinstance(dt, str | int | float):
@ -19,7 +21,6 @@ class DatetimeHandler(object):
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 {
@ -39,8 +40,8 @@ class DatetimeHandler(object):
self._dt_info_dict = self._parse_dt_info()
return self._dt_info_dict
@staticmethod
def extract_date_parts_cn(input_string: str):
@classmethod
def extract_date_parts_cn(cls, input_string: str) -> dict:
# Extending our pattern to handle every/每 as a possible value.
patterns = {
"year": r"(\d+|每)年",
@ -64,12 +65,36 @@ class DatetimeHandler(object):
extracted_data[key] = int(match.group(1))
return extracted_data
def extract_date_parts(self):
@classmethod
def extract_date_parts(cls, input_string: str) -> dict:
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!")
if not hasattr(cls, func_name):
cls.logger.warning(f"language={G_CONTEXT.language} needs to complete extract_date_parts func!")
return {}
return getattr(self, func_name)()
return getattr(cls, func_name)(input_string=input_string)
@classmethod
def format_time_by_extract_time_cn(cls, extract_time_dict: Dict[str, str], meta_data: Dict[str, str]) -> str:
cn_key_dict = {"year": "", "month": "", "day": "", "weekday": ""}
format_time_str = ""
for key, value_cn in cn_key_dict.items():
if key in extract_time_dict and key in meta_data:
value = meta_data[key]
if value_cn:
if value == "-1":
value = ""
format_time_str += f"{value}{value_cn}"
else:
format_time_str += value
return format_time_str
@classmethod
def format_time_by_extract_time(cls, extract_time_dict: Dict[str, str], meta_data: Dict[str, str]) -> str:
func_name = f"format_time_by_extract_time_{G_CONTEXT.language}"
if not hasattr(cls, func_name):
cls.logger.warning(f"language={G_CONTEXT.language} needs to complete format_time_by_extract_time func!")
return ""
return getattr(cls, func_name)(extract_time_dict, meta_data)
def datetime_format(self, dt_format: str = "%Y%m%d"):
return self._dt.strftime(dt_format)