[dev] add human_profile_setting to cli chat

This commit is contained in:
jinli.yl 2024-07-03 14:11:05 +08:00
parent 3078c9399c
commit d7dab0ca7d
18 changed files with 172 additions and 97 deletions

View file

@ -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

View file

@ -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):

View file

@ -20,6 +20,8 @@ MEMORY = "memory"
DEFAULT_SYSTEM_PROMPT = "default_system_prompt"
NOT_REFLECTED_NODES = "not_reflected_nodes"
MODIFIED_MEMORIES = "modified_memories"

View file

@ -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])

View file

@ -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

View file

@ -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)}")

View file

@ -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))

View file

@ -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] = []

View 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)

View file

@ -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)

View file

@ -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)

View file

@ -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}")

View file

@ -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)

View file

@ -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} "

View file

@ -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]}"

View file

@ -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 = {}

View file

@ -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()

View file

@ -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