mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-05 08:06:15 +00:00
[dev] add human_profile_setting to cli chat
This commit is contained in:
parent
3078c9399c
commit
d7dab0ca7d
18 changed files with 172 additions and 97 deletions
|
|
@ -7,6 +7,10 @@ memory_chat:
|
|||
memory_service: memory_chat_service
|
||||
generation_model: dashscope_generation
|
||||
human_name: 用户
|
||||
human_profile_setting:
|
||||
- name
|
||||
- gender
|
||||
- residential location
|
||||
assistant_name: AI
|
||||
memory_service:
|
||||
memory_chat_service:
|
||||
|
|
@ -50,8 +54,8 @@ worker:
|
|||
generation_model_top_k: 1
|
||||
retrieve_store_worker:
|
||||
class: memory.worker.read.retrieve_store_worker
|
||||
retrieve_obs_top_k: 100
|
||||
retrieve_ins_pf_top_k: 100
|
||||
retrieve_obs_top_k: 5
|
||||
retrieve_ins_pf_top_k: 5
|
||||
semantic_rank_worker:
|
||||
class: memory.worker.read.semantic_rank_worker
|
||||
fuse_rerank_worker:
|
||||
|
|
@ -83,8 +87,12 @@ worker:
|
|||
class: memory.worker.write.contra_repeat_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_top_k: 1
|
||||
today_obs_top_k: 30
|
||||
retrieve_top_k: 30
|
||||
contra_repeat_max_count: 50
|
||||
get_reflection_worker:
|
||||
class: memory.worker.summary.get_reflection_worker
|
||||
retrieve_top_k: 100
|
||||
reflect_obs_cnt_threshold: 32
|
||||
models:
|
||||
dashscope_generation:
|
||||
class: models.llama_index_generation_model
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import os
|
||||
import time
|
||||
from typing import List
|
||||
|
||||
import questionary
|
||||
|
||||
|
|
@ -27,19 +28,26 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
memory_service: str,
|
||||
generation_model: str,
|
||||
stream: bool = True,
|
||||
human_name: str = "human",
|
||||
assistant_name: str = "assistant",
|
||||
human_name: str = "",
|
||||
human_profile_setting: List[str] = None,
|
||||
assistant_name: str = "",
|
||||
**kwargs):
|
||||
|
||||
self._memory_service: BaseMemoryService | str = memory_service
|
||||
self._generation_model: BaseModel | str = generation_model
|
||||
self.stream: bool = stream
|
||||
self.human_name: str = human_name
|
||||
self.human_profile_setting: List[str] = human_profile_setting
|
||||
self.assistant_name: str = assistant_name
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
self._logo = char_logo("MemoryScope")
|
||||
self._prompt_handler: PromptHandler | None = None
|
||||
G_CONTEXT.meta_data.update({
|
||||
"human_name": human_name,
|
||||
"assistant_name": assistant_name,
|
||||
"human_profile_setting": human_profile_setting,
|
||||
})
|
||||
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
|
|
@ -66,7 +74,8 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
self._generation_model = G_CONTEXT.model_dict[self._generation_model]
|
||||
return self._generation_model
|
||||
|
||||
def get_system_prompt(self) -> Message:
|
||||
@property
|
||||
def system_prompt_with_memory(self) -> Message:
|
||||
system_prompt = self.prompt_handler.system_prompt
|
||||
|
||||
memories: str = self.memory_service.read_memory()
|
||||
|
|
@ -77,13 +86,9 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt)
|
||||
|
||||
def chat_with_memory(self, query: str) -> ModelResponse | ModelResponseGen:
|
||||
new_message: Message = Message(role=MessageRoleEnum.USER.value,
|
||||
role_name=self.human_name,
|
||||
content=query)
|
||||
|
||||
new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=self.human_name, content=query)
|
||||
self.memory_service.add_messages(new_message)
|
||||
system_message: Message = self.get_system_prompt()
|
||||
return self.generation_model.call(messages=[system_message, new_message], stream=self.stream)
|
||||
return self.generation_model.call(messages=[self.system_prompt_with_memory, new_message], stream=self.stream)
|
||||
|
||||
@staticmethod
|
||||
def parse_query_command(query: str):
|
||||
|
|
|
|||
|
|
@ -20,6 +20,8 @@ MEMORY = "memory"
|
|||
|
||||
DEFAULT_SYSTEM_PROMPT = "default_system_prompt"
|
||||
|
||||
NOT_REFLECTED_NODES = "not_reflected_nodes"
|
||||
|
||||
|
||||
MODIFIED_MEMORIES = "modified_memories"
|
||||
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@ class BaseWorker(metaclass=ABCMeta):
|
|||
self.logger: Logger = Logger.get_logger()
|
||||
|
||||
@staticmethod
|
||||
def _async_run(fn_list, *args, **kwargs):
|
||||
def async_run(fn_list, *args, **kwargs):
|
||||
async def async_gather():
|
||||
return await asyncio.gather(*[fn(*args, **kwargs) for fn in fn_list])
|
||||
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ from abc import ABCMeta
|
|||
from typing import List, Dict
|
||||
|
||||
from memory_scope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS
|
||||
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memory_scope.memory.worker.base_worker import BaseWorker
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.scheme.message import Message
|
||||
|
|
@ -77,15 +76,13 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
@property
|
||||
def user_name(self) -> str:
|
||||
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
|
||||
self._user_name = G_CONTEXT.meta_data["human_name"]
|
||||
return self._user_name
|
||||
|
||||
@property
|
||||
def target_name(self) -> str:
|
||||
if self._target_name is None:
|
||||
message = [x for x in self.messages if x.role == MessageRoleEnum.USER.value][-1]
|
||||
self._target_name = message.role_name
|
||||
self._target_name = G_CONTEXT.meta_data["assistant_name"]
|
||||
return self._target_name
|
||||
|
||||
@property
|
||||
|
|
|
|||
|
|
@ -35,7 +35,7 @@ class RetrieveStoreWorker(MemoryBaseWorker):
|
|||
query, _ = self.get_context(QUERY_WITH_TS)
|
||||
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.async_run(fn_list=fn_list, query=query):
|
||||
if result:
|
||||
memory_node_list.extend(result)
|
||||
self.logger.info(f"memory_node_list.size={len(memory_node_list)}")
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ class SetQueryWorker(MemoryBaseWorker):
|
|||
query = self.chat_kwargs["query"]
|
||||
query_timestamp = int(datetime.datetime.now().timestamp())
|
||||
else:
|
||||
query = self.messages[-1].content
|
||||
query_timestamp = self.messages[-1].time_created
|
||||
query = self.chat_messages[-1].content
|
||||
query_timestamp = self.chat_messages[-1].time_created
|
||||
|
||||
self.set_context(QUERY_WITH_TS, (query, query_timestamp))
|
||||
|
|
|
|||
|
|
@ -1,47 +1,39 @@
|
|||
from typing import List
|
||||
|
||||
from memory_scope.utils.response_text_parser import ResponseTextParser
|
||||
from memory_scope.constants.common_constants import (
|
||||
NEW_OBS_NODES,
|
||||
NOT_REFLECTED_OBS_NODES,
|
||||
INSIGHT_NODES,
|
||||
NEW_INSIGHT_KEYS,
|
||||
NOT_REFLECTED_MERGE_NODES,
|
||||
)
|
||||
from memory_scope.constants.language_constants import COLON_WORD, COMMA_WORD
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.constants.common_constants import NOT_REFLECTED_NODES
|
||||
from memory_scope.constants.language_constants import COMMA_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.utils.response_text_parser import ResponseTextParser
|
||||
from memory_scope.utils.timer import timer
|
||||
from memory_scope.utils.tool_functions import prompt_to_msg
|
||||
|
||||
|
||||
class GetReflectionWorker(MemoryBaseWorker):
|
||||
@timer
|
||||
def retrieve_not_reflected_memory(self) -> List[MemoryNode]:
|
||||
filter_dict = {
|
||||
"user_name": self.user_name,
|
||||
"target_name": self.target_name,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"obs_reflected": False,
|
||||
}
|
||||
return self.vector_store.retrieve(query=" ", top_k=self.retrieve_top_k, filter_dict=filter_dict)
|
||||
|
||||
def _run(self):
|
||||
# 过滤得到 not_reflected_merge_nodes
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_context(
|
||||
NOT_REFLECTED_OBS_NODES
|
||||
)
|
||||
not_reflected_merge_nodes: List[MemoryNode] = []
|
||||
if new_obs_nodes:
|
||||
not_reflected_merge_nodes.extend(new_obs_nodes)
|
||||
if not_reflected_nodes:
|
||||
not_reflected_merge_nodes.extend(not_reflected_nodes)
|
||||
not_reflected_merge_nodes = [
|
||||
node
|
||||
for node in not_reflected_merge_nodes
|
||||
if not node.obs_reflected
|
||||
]
|
||||
not_reflected_nodes: List[MemoryNode] = self.retrieve_not_reflected_memory()
|
||||
|
||||
# count
|
||||
not_reflected_count = len(not_reflected_merge_nodes)
|
||||
not_reflected_count = len(not_reflected_nodes)
|
||||
if not_reflected_count <= self.reflect_obs_cnt_threshold:
|
||||
self.logger.info(
|
||||
f"not_reflected_count={not_reflected_count} is not enough, stop reflect."
|
||||
)
|
||||
self.logger.info(f"not_reflected_count={not_reflected_count} is not enough, stop.")
|
||||
return
|
||||
|
||||
# save context
|
||||
self.set_context(NOT_REFLECTED_MERGE_NODES, not_reflected_merge_nodes)
|
||||
self.set_context(NOT_REFLECTED_NODES, not_reflected_nodes)
|
||||
|
||||
# get profile_keys
|
||||
exist_keys: List[str] = []
|
||||
|
|
|
|||
68
memory_scope/memory/worker/summary/load_memory_worker.py
Normal file
68
memory_scope/memory/worker/summary/load_memory_worker.py
Normal file
|
|
@ -0,0 +1,68 @@
|
|||
from typing import List, Dict
|
||||
|
||||
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.utils.timer import timer
|
||||
from memory_scope.utils.global_context import G_CONTEXT
|
||||
|
||||
class LoadMemoryWorker(MemoryBaseWorker):
|
||||
|
||||
@timer
|
||||
async def retrieve_not_reflected_memory(self, query: str) -> List[MemoryNode]:
|
||||
filter_dict = {
|
||||
"user_name": self.user_name,
|
||||
"target_name": self.target_name,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"obs_reflected": False,
|
||||
}
|
||||
return await self.vector_store.async_retrieve(query=query,
|
||||
top_k=self.retrieve_not_reflected_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
||||
@timer
|
||||
async def retrieve_not_updated_memory(self, query: str) -> List[MemoryNode]:
|
||||
filter_dict = {
|
||||
"user_name": self.user_name,
|
||||
"target_name": self.target_name,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"obs_updated": False,
|
||||
}
|
||||
return await self.vector_store.async_retrieve(query=query,
|
||||
top_k=self.retrieve_not_updated_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
||||
@timer
|
||||
async def retrieve_profiles(self, query: str) -> List[MemoryNode]:
|
||||
filter_dict = {
|
||||
"user_name": self.user_name,
|
||||
"target_name": self.target_name,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [MemoryTypeEnum.PROFILE.value, MemoryTypeEnum.PROFILE_CUSTOMIZED.value],
|
||||
}
|
||||
retrieve_nodes = await self.vector_store.async_retrieve(query=query,
|
||||
top_k=self.retrieve_profiles_top_k,
|
||||
filter_dict=filter_dict)
|
||||
nodes: List[MemoryNode] = []
|
||||
human_profile_setting = G_CONTEXT.meta_data.get("human_profile_setting", [])
|
||||
for attr_key in human_profile_setting:
|
||||
|
||||
|
||||
return nodes
|
||||
|
||||
def _run(self):
|
||||
mock_query = "_"
|
||||
fn_list = [
|
||||
self.retrieve_not_reflected_memory,
|
||||
self.retrieve_not_updated_memory,
|
||||
self.retrieve_profiles,
|
||||
]
|
||||
memory_node_dict: Dict[str, MemoryNode] = {}
|
||||
for nodes in self.async_run(fn_list, query=mock_query):
|
||||
assert isinstance(nodes[0], MemoryNode)
|
||||
memory_node_dict.update({n.memory_id: n for n in nodes})
|
||||
|
||||
memory_node_list = sorted(memory_node_dict.values(), key=lambda x: x.memory_id)
|
||||
|
|
@ -27,7 +27,7 @@ class ContraRepeatWorker(MemoryBaseWorker):
|
|||
"target_name": self.target_name,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"obs_dt": dt_handler.datetime_format(),
|
||||
"dt": dt_handler.datetime_format(),
|
||||
}
|
||||
return self.vector_store.retrieve(query=message.content, top_k=self.today_obs_top_k, filter_dict=filter_dict)
|
||||
|
||||
|
|
@ -55,11 +55,11 @@ class ContraRepeatWorker(MemoryBaseWorker):
|
|||
for i, n in enumerate(all_obs_nodes):
|
||||
user_query_list.append(f"{i + 1} {n.content}")
|
||||
|
||||
system_prompt = self.prompt_config.contra_repeat_system.format(num_obs=len(user_query_list),
|
||||
user_name=self.user_id)
|
||||
few_shot = self.prompt_config.contra_repeat_few_shot.format(user_name=self.user_id)
|
||||
user_query = self.prompt_config.contra_repeat_user_query.format(user_query="\n".join(user_query_list),
|
||||
system_prompt = self.prompt_handler.contra_repeat_system.format(num_obs=len(user_query_list),
|
||||
user_name=self.user_id)
|
||||
few_shot = self.prompt_handler.contra_repeat_few_shot.format(user_name=self.user_id)
|
||||
user_query = self.prompt_handler.contra_repeat_user_query.format(user_query="\n".join(user_query_list),
|
||||
user_name=self.user_id)
|
||||
contra_repeat_message = self.prompt_to_msg(system_prompt=system_prompt,
|
||||
few_shot=few_shot,
|
||||
user_query=user_query)
|
||||
|
|
|
|||
|
|
@ -33,8 +33,8 @@ class GetObservationWithTimeWorker(GetObservationWorker):
|
|||
|
||||
system_prompt = self.prompt_handler.get_observation_with_time_system.format(num_obs=len(user_query_list),
|
||||
user_name=self.target_name)
|
||||
few_shot = self.prompt_config.get_observation_with_time_few_shot.format(user_name=self.target_name)
|
||||
user_query = self.prompt_config.get_observation_with_time_user_query.format(
|
||||
few_shot = self.prompt_handler.get_observation_with_time_few_shot.format(user_name=self.target_name)
|
||||
user_query = self.prompt_handler.get_observation_with_time_user_query.format(
|
||||
user_query="\n".join(user_query_list),
|
||||
user_name=self.target_name)
|
||||
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ class GetObservationWorker(MemoryBaseWorker):
|
|||
meta_data = {
|
||||
MemoryTypeEnum.CONVERSATION.value: message.content,
|
||||
TIME_INFER: time_infer,
|
||||
"keywords": keywords,
|
||||
**dt_handler.dt_info_dict,
|
||||
}
|
||||
|
||||
|
|
@ -34,11 +35,8 @@ class GetObservationWorker(MemoryBaseWorker):
|
|||
memory_type=MemoryTypeEnum.OBSERVATION.value,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
timestamp=message.time_created,
|
||||
obs_dt=dt_handler.datetime_format(),
|
||||
obs_reflected=False,
|
||||
obs_updated=False,
|
||||
obs_keyword=keywords)
|
||||
node.gen_memory_id()
|
||||
obs_updated=False)
|
||||
return node
|
||||
|
||||
def build_prompt(self) -> List[Message]:
|
||||
|
|
@ -61,9 +59,9 @@ class GetObservationWorker(MemoryBaseWorker):
|
|||
|
||||
system_prompt = self.prompt_handler.get_observation_system.format(num_obs=len(user_query_list),
|
||||
user_name=self.target_name)
|
||||
few_shot = self.prompt_config.get_observation_few_shot.format(user_name=self.target_name)
|
||||
user_query = self.prompt_config.get_observation_user_query.format(user_query="\n".join(user_query_list),
|
||||
user_name=self.target_name)
|
||||
few_shot = self.prompt_handler.get_observation_few_shot.format(user_name=self.target_name)
|
||||
user_query = self.prompt_handler.get_observation_user_query.format(user_query="\n".join(user_query_list),
|
||||
user_name=self.target_name)
|
||||
|
||||
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}")
|
||||
|
|
|
|||
|
|
@ -23,14 +23,12 @@ class StoreMemoryWorker(MemoryBaseWorker):
|
|||
return
|
||||
|
||||
dt_handler = DatetimeHandler()
|
||||
node = MemoryNode(
|
||||
user_name=self.user_name,
|
||||
target_name=self.target_name,
|
||||
content=query,
|
||||
memory_type=MemoryTypeEnum.OBSERVATION.value,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
timestamp=dt_handler.timestamp,
|
||||
obs_dt=dt_handler.datetime_format(),
|
||||
obs_reflected=False,
|
||||
obs_updated=False)
|
||||
node = MemoryNode(user_name=self.user_name,
|
||||
target_name=self.target_name,
|
||||
content=query,
|
||||
memory_type=MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
timestamp=dt_handler.timestamp,
|
||||
obs_reflected=False,
|
||||
obs_updated=False)
|
||||
self.vector_store.update(node)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
import inspect
|
||||
import time
|
||||
import dashscope
|
||||
from abc import abstractmethod, ABCMeta
|
||||
from typing import Any
|
||||
|
||||
|
|
@ -23,6 +22,7 @@ class BaseModel(metaclass=ABCMeta):
|
|||
max_retries: int = 3,
|
||||
retry_interval: float = 1.0,
|
||||
kwargs_filter: bool = True,
|
||||
raise_exception: bool = True,
|
||||
**kwargs):
|
||||
|
||||
self.model_name: str = model_name
|
||||
|
|
@ -31,6 +31,7 @@ class BaseModel(metaclass=ABCMeta):
|
|||
self.max_retries: int = max_retries
|
||||
self.retry_interval: float = retry_interval
|
||||
self.kwargs_filter: bool = kwargs_filter
|
||||
self.raise_exception: bool = raise_exception
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
self.data = {}
|
||||
|
|
@ -85,12 +86,13 @@ class BaseModel(metaclass=ABCMeta):
|
|||
with Timer(self.__class__.__name__, log_time=False) as t:
|
||||
self.before_call(stream=stream, **kwargs)
|
||||
for i in range(self.max_retries):
|
||||
try:
|
||||
if self.raise_exception:
|
||||
model_response = self._call(stream=stream, **kwargs)
|
||||
except dashscope.common.error.AuthenticationError as e:
|
||||
raise e
|
||||
except Exception as e:
|
||||
model_response = ModelResponse(m_type=self.m_type, status=False, details=e.args)
|
||||
else:
|
||||
try:
|
||||
model_response = self._call(stream=stream, **kwargs)
|
||||
except Exception as e:
|
||||
model_response = ModelResponse(m_type=self.m_type, status=False, details=e.args)
|
||||
|
||||
if isinstance(model_response, ModelResponse) and not model_response.status:
|
||||
self.logger.warning(f"call model={self.model_name} failed! cost={t.cost_str} retry_cnt={i} "
|
||||
|
|
@ -114,10 +116,13 @@ class BaseModel(metaclass=ABCMeta):
|
|||
with Timer(self.__class__.__name__, log_time=False) as t:
|
||||
self.before_call(**kwargs)
|
||||
for i in range(self.max_retries):
|
||||
try:
|
||||
model_response = await self._async_call(**kwargs)
|
||||
except Exception as e:
|
||||
model_response = ModelResponse(status=False, details=e.args)
|
||||
if self.raise_exception:
|
||||
model_response = self._async_call(**kwargs)
|
||||
else:
|
||||
try:
|
||||
model_response = self._async_call(**kwargs)
|
||||
except Exception as e:
|
||||
model_response = ModelResponse(m_type=self.m_type, status=False, details=e.args)
|
||||
|
||||
if not model_response.status:
|
||||
self.logger.warning(f"async_call model={self.model_name} failed! cost={t.cost_str} retry_cnt={i} "
|
||||
|
|
|
|||
|
|
@ -31,21 +31,16 @@ class MemoryNode(BaseModel):
|
|||
|
||||
timestamp: int = Field(int(datetime.datetime.now().timestamp()), description="timestamp of the memory node")
|
||||
|
||||
obs_dt: str = Field("", description="dt of the observation")
|
||||
dt: str = Field("", description="dt of the memory node")
|
||||
|
||||
obs_reflected: bool = Field(False, description="if the observation is reflected")
|
||||
|
||||
obs_updated: bool = Field(False, description="if the observation has updated user profile or insight")
|
||||
|
||||
obs_keyword: str = Field("", description="keywords of the content")
|
||||
|
||||
insight_key: str = Field("", description="insight_key")
|
||||
|
||||
insight_value: str = Field("", description="insight_value")
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.gen_memory_id()
|
||||
self.memory_id = f"{self.user_name}_{self.target_name}_{self.timestamp}_{md5_hash(self.content)[:8]}"
|
||||
self.dt = datetime.datetime.fromtimestamp(self.timestamp).strftime("%Y%m%d")
|
||||
|
||||
@property
|
||||
def node_keys(self):
|
||||
|
|
@ -54,5 +49,3 @@ class MemoryNode(BaseModel):
|
|||
def __getitem__(self, key: str):
|
||||
return self.model_dump().get(key)
|
||||
|
||||
def gen_memory_id(self):
|
||||
self.memory_id = f"{self.user_name}_{self.target_name}_{self.timestamp}_{md5_hash(self.content)[:8]}"
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from llama_index.vector_stores.elasticsearch import ElasticsearchStore, AsyncDen
|
|||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.storage.base_vector_store import BaseVectorStore
|
||||
from memory_scope.utils.logger import Logger
|
||||
|
||||
|
||||
class _ElasticsearchStore(ElasticsearchStore):
|
||||
|
|
@ -83,6 +84,7 @@ class LlamaIndexElasticSearchStore(BaseVectorStore):
|
|||
**kwargs)
|
||||
self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store,
|
||||
embed_model=self.embedding_model.model)
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
def retrieve(self,
|
||||
query: str,
|
||||
|
|
@ -100,6 +102,8 @@ class LlamaIndexElasticSearchStore(BaseVectorStore):
|
|||
query: str,
|
||||
top_k: int,
|
||||
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
|
||||
self.logger.info(f"query={query} top_k={top_k} filter_dict={filter_dict}")
|
||||
|
||||
if filter_dict is None:
|
||||
filter_dict = {}
|
||||
|
||||
|
|
|
|||
|
|
@ -23,5 +23,7 @@ class GlobalContext(object):
|
|||
self.thread_pool: ThreadPoolExecutor | None = None
|
||||
self.language: LanguageEnum = LanguageEnum.EN
|
||||
|
||||
self.meta_data: Dict[str, Any] = {}
|
||||
|
||||
|
||||
G_CONTEXT = GlobalContext()
|
||||
|
|
|
|||
|
|
@ -6,10 +6,13 @@ from copy import deepcopy
|
|||
from importlib import import_module
|
||||
|
||||
import pyfiglet
|
||||
from termcolor import colored, COLORS
|
||||
from termcolor import colored
|
||||
|
||||
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
|
||||
|
||||
ALL_COLORS = ["red", "green", "yellow", "blue", "magenta", "cyan", "light_grey", "light_red", "light_green",
|
||||
"light_yellow", "light_blue", "light_magenta", "light_cyan", "white"]
|
||||
|
||||
|
||||
def underscore_to_camelcase(name: str, is_first_title: bool = True):
|
||||
name_split = name.split("_")
|
||||
|
|
@ -69,7 +72,7 @@ def char_logo(words: str, seed: int = time.time_ns(), color=None):
|
|||
font = pyfiglet.Figlet()
|
||||
rendered_text = font.renderText(words)
|
||||
colored_lines = []
|
||||
all_colors = list(COLORS.keys())
|
||||
all_colors = ALL_COLORS.copy()
|
||||
random.seed = seed
|
||||
for line in rendered_text.splitlines():
|
||||
line_color = color
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue