mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-05 08:06:15 +00:00
[dev] reformat llm params
This commit is contained in:
parent
b151da319e
commit
c2406e2eba
8 changed files with 75 additions and 43 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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]}"
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue