[dev] rename human default name

This commit is contained in:
jinli.yl 2024-07-08 15:36:15 +08:00
parent 8e756fdd98
commit f947825cba
15 changed files with 145 additions and 136 deletions

View file

@ -4,6 +4,7 @@ import time
import questionary
from memory_scope.chat.base_memory_chat import BaseMemoryChat
from memory_scope.constants.language_constants import DEFAULT_HUMAN_NAME
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
from memory_scope.memory.service.base_memory_service import BaseMemoryService
from memory_scope.models.base_model import BaseModel
@ -27,7 +28,7 @@ class CliMemoryChat(BaseMemoryChat):
memory_service: str,
generation_model: str,
stream: bool = True,
human_name: str = "用户",
human_name: str = DEFAULT_HUMAN_NAME[G_CONTEXT.language],
assistant_name: str = "AI",
**kwargs):
@ -192,9 +193,7 @@ class CliMemoryChat(BaseMemoryChat):
except Exception as e:
import traceback
traceback.print_exc()
line = f"An exception occurred when running cli memory chat. args={e.args}."
questionary.print(line)
self.logger.exception(line)
self.logger.exception(f"An exception occurred when running cli memory chat. args={e.args}.")
continue
questionary.print(f"A memory writing thread is still running, please be patient and wait!")

View file

@ -37,6 +37,8 @@ NEW_OBS_WITH_TIME_NODES = "new_obs_with_time_nodes"
INSIGHT_NODES = "insight_nodes"
TODAY_NODES = "today_nodes"
MERGE_OBS_NODES = "merge_obs_nodes"
NEW_INSIGHT_NODES = "new_insight_nodes"

View file

@ -74,3 +74,8 @@ COMMA_WORD = {
LanguageEnum.CN: "",
LanguageEnum.EN: ","
}
DEFAULT_HUMAN_NAME = {
LanguageEnum.CN: "用户",
LanguageEnum.EN: "user"
}

View file

@ -1,5 +1,5 @@
from abc import ABCMeta
from typing import List, Dict
from typing import List, Dict, Set
from memory_scope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS
from memory_scope.memory.worker.base_worker import BaseWorker
@ -70,18 +70,21 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
self._memory_store = G_CONTEXT.memory_store
return self._memory_store
def get_memories(self, key: str) -> List[MemoryNode]:
def get_memories(self, keys: str | List[str]) -> List[MemoryNode]:
memories: List[MemoryNode] = []
memory_ids: List[str] = self.get_context(key)
if memory_ids:
memories.extend([self._contex_memory_dict[x] for x in memory_ids])
if isinstance(keys, str):
keys = [keys]
for key in keys:
memory_ids: List[str] = self.get_context(key)
if memory_ids:
memories.extend([self._contex_memory_dict[x] for x in memory_ids])
return memories
def set_memories(self, key: str, nodes: List[MemoryNode] | MemoryNode):
if not nodes:
return
if isinstance(nodes, MemoryNode):
def set_memories(self, key: str, nodes: MemoryNode | List[MemoryNode]):
if nodes is None:
nodes = []
elif isinstance(nodes, MemoryNode):
nodes = [nodes]
for node in nodes:
@ -90,6 +93,23 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
self._contex_memory_dict[node.memory_id] = node
self.set_context(key, [n.memory_id for n in nodes])
def save_memories(self, keys: str | List[str] = None):
if keys is None:
self.memory_store.update_memories(list(self._contex_memory_dict.values()))
self._contex_memory_dict.clear()
return
if isinstance(keys, str):
keys = [keys]
ids: Set[str] = Set[str]()
for key in keys:
t_ids: List[str] = self.get_context(key)
if t_ids:
ids.update(t_ids)
nodes = [self._contex_memory_dict.pop(_) for _ in ids]
self.memory_store.update_memories(nodes)
@property
def monitor(self) -> BaseMonitor:
if self._monitor is None:

View file

@ -6,51 +6,38 @@ 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.datetime_handler import DatetimeHandler
from memory_scope.utils.timer import timer
class PrintMemoryWorker(MemoryBaseWorker):
@timer
def retrieve_expired_memory(self, query: str):
filter_dict = {
"user_name": self.user_name,
"target_name": self.target_name,
"status": MemoryNodeStatus.EXPIRED.value,
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
}
return self.memory_store.retrieve_memories(query=query,
top_k=self.retrieve_expired_top_k,
filter_dict=filter_dict)
def _run(self):
expired_memories: List[MemoryNode] = self.retrieve_expired_memory(query="_")
memory_node_list: List[MemoryNode] = self.get_memories(RETRIEVE_MEMORY_NODES)
memory_node_list = sorted(memory_node_list, key=lambda x: x.timestamp, reverse=True)
expired_content_list: List[str] = []
obs_content_list: List[str] = []
insight_content_list: List[str] = []
expired_content_list: List[str] = []
i = 0
j = 0
k = 0
for node in memory_node_list:
if MemoryNodeStatus(node.status) is MemoryNodeStatus.EXPIRED:
i += 1
line = f" {i} {node.content}"
expired_content_list.append(line)
if MemoryTypeEnum(node.memory_type) in [MemoryTypeEnum.OBSERVATION, MemoryTypeEnum.OBS_CUSTOMIZED]:
i += 1
j += 1
dt_handler = DatetimeHandler(node.timestamp)
dt = dt_handler.datetime_format("%Y%m%d %H:%M:%S")
line = f" {i} {dt} {node.content}"
line = f" {j} {dt} {node.content}"
obs_content_list.append(line)
elif MemoryTypeEnum(node.memory_type) in [MemoryTypeEnum.INSIGHT, ]:
j += 1
line = f" {j} {node.content}"
elif MemoryTypeEnum(node.memory_type) is MemoryTypeEnum.INSIGHT:
k += 1
line = f" {k} {node.content}"
insight_content_list.append(line)
for i, node in enumerate(expired_memories):
line = f" {j} {node.content}"
expired_content_list.append(line)
obs_content = "\n".join(obs_content_list)
insight_content = "\n".join(insight_content_list)
expired_content = "\n".join(expired_content_list)

View file

@ -5,11 +5,16 @@ 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
class RetrieveStoreWorker(MemoryBaseWorker):
class RetrieveMemoryWorker(MemoryBaseWorker):
@timer
async def retrieve_from_observation(self, query: str) -> List[MemoryNode]:
if not self.retrieve_obs_top_k:
return []
filter_dict = {
"user_name": self.user_name,
"target_name": self.target_name,
@ -20,7 +25,11 @@ class RetrieveStoreWorker(MemoryBaseWorker):
top_k=self.retrieve_obs_top_k,
filter_dict=filter_dict)
@timer
async def retrieve_from_insight_and_profile(self, query: str) -> List[MemoryNode]:
if not self.retrieve_ins_pf_top_k:
return []
filter_dict = {
"user_name": self.user_name,
"target_name": self.target_name,
@ -31,10 +40,26 @@ class RetrieveStoreWorker(MemoryBaseWorker):
top_k=self.retrieve_ins_pf_top_k,
filter_dict=filter_dict)
@timer
async def retrieve_expired_memory(self, query: str) -> List[MemoryNode]:
if not self.retrieve_expired_top_k:
return []
filter_dict = {
"user_name": self.user_name,
"target_name": self.target_name,
"status": MemoryNodeStatus.EXPIRED.value,
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
}
return await self.memory_store.a_retrieve_memories(query=query,
top_k=self.retrieve_expired_top_k,
filter_dict=filter_dict)
def _run(self):
query, _ = self.get_context(QUERY_WITH_TS)
self.submit_async_task(self.retrieve_from_observation, query=query)
self.submit_async_task(self.retrieve_from_insight_and_profile, query=query)
self.submit_async_task(self.retrieve_expired_memory, query=query)
memory_node_list: List[MemoryNode] = []
for result in self.gather_async_result():
@ -44,5 +69,6 @@ class RetrieveStoreWorker(MemoryBaseWorker):
memory_node_list = sorted(memory_node_list, key=lambda x: x.score_similar, reverse=True)
for node in memory_node_list:
self.logger.info(f"recall_stage: content={node.content} score={node.score_similar} type={node.memory_type}")
self.logger.info(f"recall_stage: content={node.content} score={node.score_similar} type={node.memory_type}"
f"status={node.status}")
self.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list)

View file

@ -25,8 +25,8 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
status=MemoryNodeStatus.NEW.value)
def _run(self):
not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_NODES)
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
not_reflected_nodes: List[MemoryNode] = self.get_memories(NOT_REFLECTED_NODES)
insight_nodes: List[MemoryNode] = self.get_memories(INSIGHT_NODES)
# count
not_reflected_count = len(not_reflected_nodes)

View file

@ -1,8 +1,6 @@
from typing import List, Dict
from memory_scope.constants.common_constants import (
NOT_UPDATED_NODES, MERGE_OBS_NODES,
)
from memory_scope.constants.common_constants import NOT_UPDATED_NODES, MERGE_OBS_NODES
from memory_scope.constants.language_constants import NONE_WORD, INCLUDED_WORD, CONTRADICTORY_WORD
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
@ -27,7 +25,7 @@ class LongContraRepeatWorker(MemoryBaseWorker):
return node, [n for n in retrieve_nodes if n.score_similar >= self.long_contra_repeat_threshold]
def _run(self):
not_updated_nodes: List[MemoryNode] = self.get_context(NOT_UPDATED_NODES)
not_updated_nodes: List[MemoryNode] = self.get_memories(NOT_UPDATED_NODES)
for node in not_updated_nodes:
self.submit_async_task(fn=self.retrieve_similar_content, node=node)
@ -103,10 +101,13 @@ class LongContraRepeatWorker(MemoryBaseWorker):
node.status = MemoryNodeStatus.EXPIRED.value
else:
node.content = content
node.status = MemoryNodeStatus.CONTENT_MODIFIED.value
elif status == self.get_language_value(INCLUDED_WORD):
node.status = MemoryNodeStatus.EXPIRED.value
merge_obs_nodes.append(node)
self.logger.info(f"after_long_contra_repeat: {node.content} {node.status}")
# save context
self.set_context(MERGE_OBS_NODES, merge_obs_nodes)
self.set_memories(MERGE_OBS_NODES, merge_obs_nodes)

View file

@ -1,40 +0,0 @@
from typing import List, Dict
from memory_scope.constants.common_constants import NOT_REFLECTED_NODES, INSIGHT_NODES, MERGE_OBS_NODES, \
NOT_UPDATED_NODES
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
from memory_scope.scheme.memory_node import MemoryNode
class SummaryCollectWorker(MemoryBaseWorker):
def _run(self):
update_memories: Dict[str, MemoryNode] = {}
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
if insight_nodes:
update_memories.update({n.memory_id: n for n in insight_nodes})
not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_NODES)
if not_reflected_nodes:
update_memories.update({n.memory_id: n for n in not_reflected_nodes})
not_updated_nodes: List[MemoryNode] = self.get_context(NOT_UPDATED_NODES)
if not_updated_nodes:
for node in not_updated_nodes:
if node.memory_id in update_memories:
keys = [
INSIGHT_NODES,
MERGE_OBS_NODES,
NOT_UPDATED_NODES,
NOT_REFLECTED_NODES,
]
memory_nodes: List[MemoryNode] = []
for key in keys:
memory_nodes.extend(self.get_context(key))
self.memory_store.update_memories(update_memories)

View file

@ -100,9 +100,9 @@ class UpdateInsightWorker(MemoryBaseWorker):
return insight_node
def _run(self):
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
not_updated_nodes: List[MemoryNode] = self.get_context(NOT_UPDATED_NODES)
not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_NODES)
insight_nodes: List[MemoryNode] = self.get_memories(INSIGHT_NODES)
not_updated_nodes: List[MemoryNode] = self.get_memories(NOT_UPDATED_NODES)
not_reflected_nodes: List[MemoryNode] = self.get_memories(NOT_REFLECTED_NODES)
if not insight_nodes:
self.logger.warning("insight_nodes is empty, stop.")
@ -136,3 +136,6 @@ class UpdateInsightWorker(MemoryBaseWorker):
for node in not_updated_nodes:
node.obs_updated = True
for node in not_reflected_nodes:
node.obs_updated = True

View file

@ -1,51 +1,22 @@
from typing import List
from memory_scope.constants.common_constants import NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, MERGE_OBS_NODES
from memory_scope.constants.common_constants import NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, MERGE_OBS_NODES, TODAY_NODES
from memory_scope.constants.language_constants import NONE_WORD, CONTRADICTORY_WORD, INCLUDED_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.scheme.message import Message
from memory_scope.utils.datetime_handler import DatetimeHandler
from memory_scope.utils.response_text_parser import ResponseTextParser
from memory_scope.utils.timer import timer
class ContraRepeatWorker(MemoryBaseWorker):
@timer
def retrieve_today_memory(self) -> List[MemoryNode]:
if not self.chat_messages:
self.logger.warning("chat_messages is empty!")
return []
message: Message = self.chat_messages[-1]
dt_handler = DatetimeHandler(message.time_created)
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],
"dt": dt_handler.datetime_format(),
}
return self.memory_store.retrieve_memories(query=message.content, top_k=self.today_obs_top_k, filter_dict=filter_dict)
def _run(self):
all_obs_nodes: List[MemoryNode] = []
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
if new_obs_nodes:
all_obs_nodes.extend(new_obs_nodes)
new_obs_with_time_nodes: List[MemoryNode] = self.get_context(NEW_OBS_WITH_TIME_NODES)
if new_obs_with_time_nodes:
all_obs_nodes.extend(new_obs_with_time_nodes)
all_obs_nodes: List[MemoryNode] = self.get_memories([NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES])
if not all_obs_nodes:
self.logger.info("all_obs_nodes is empty!")
self.continue_run = False
return
today_obs_nodes: List[MemoryNode] = self.retrieve_today_memory()
today_obs_nodes: List[MemoryNode] = self.get_memories(TODAY_NODES)
if today_obs_nodes:
all_obs_nodes.extend(today_obs_nodes)
all_obs_nodes = sorted(all_obs_nodes, key=lambda x: x.timestamp, reverse=True)[:self.contra_repeat_max_count]
@ -109,11 +80,11 @@ class ContraRepeatWorker(MemoryBaseWorker):
node: MemoryNode = all_obs_nodes[idx]
if keep_flag != self.get_language_value(NONE_WORD):
node.memory_node.status = MemoryNodeStatus.EXPIRED.value
node.status = MemoryNodeStatus.EXPIRED.value
merge_obs_nodes.append(node)
# forbid keyword
self.logger.info(f"contra_repeat stage: {node.content} {node.status}")
# save context
self.set_context(MERGE_OBS_NODES, merge_obs_nodes)
self.set_memories(MERGE_OBS_NODES, merge_obs_nodes)

View file

@ -38,4 +38,4 @@ class GetObservationWithTimeWorker(GetObservationWorker):
return obtain_obs_message
def save(self, new_obs_nodes: List[MemoryNode]):
self.set_context(NEW_OBS_WITH_TIME_NODES, new_obs_nodes)
self.set_memories(NEW_OBS_WITH_TIME_NODES, new_obs_nodes)

View file

@ -62,7 +62,7 @@ class GetObservationWorker(MemoryBaseWorker):
return obtain_obs_message
def save(self, new_obs_nodes: List[MemoryNode]):
self.set_context(NEW_OBS_NODES, new_obs_nodes)
self.set_memories(NEW_OBS_NODES, new_obs_nodes)
def _run(self):
obtain_obs_message = self.build_prompt()

View file

@ -1,10 +1,12 @@
from typing import List
from memory_scope.constants.common_constants import NOT_REFLECTED_NODES, NOT_UPDATED_NODES, INSIGHT_NODES
from memory_scope.constants.common_constants import NOT_REFLECTED_NODES, NOT_UPDATED_NODES, INSIGHT_NODES, TODAY_NODES
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.scheme.message import Message
from memory_scope.utils.datetime_handler import DatetimeHandler
from memory_scope.utils.timer import timer
@ -12,6 +14,9 @@ class LoadMemoryWorker(MemoryBaseWorker):
@timer
async def retrieve_not_reflected_memory(self, query: str):
if not self.retrieve_not_reflected_top_k:
return
filter_dict = {
"user_name": self.user_name,
"target_name": self.target_name,
@ -22,10 +27,13 @@ class LoadMemoryWorker(MemoryBaseWorker):
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query,
top_k=self.retrieve_not_reflected_top_k,
filter_dict=filter_dict)
self.set_context(NOT_REFLECTED_NODES, nodes)
self.set_memories(NOT_REFLECTED_NODES, nodes)
@timer
async def retrieve_not_updated_memory(self, query: str):
if not self.retrieve_not_updated_top_k:
return
filter_dict = {
"user_name": self.user_name,
"target_name": self.target_name,
@ -36,10 +44,13 @@ class LoadMemoryWorker(MemoryBaseWorker):
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query,
top_k=self.retrieve_not_updated_top_k,
filter_dict=filter_dict)
self.set_context(NOT_UPDATED_NODES, nodes)
self.set_memories(NOT_UPDATED_NODES, nodes)
@timer
async def retrieve_insight_memory(self, query: str):
if not self.retrieve_insight_top_k:
return
filter_dict = {
"user_name": self.user_name,
"target_name": self.target_name,
@ -49,12 +60,36 @@ class LoadMemoryWorker(MemoryBaseWorker):
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query,
top_k=self.retrieve_insight_top_k,
filter_dict=filter_dict)
self.set_context(INSIGHT_NODES, nodes)
self.set_memories(INSIGHT_NODES, nodes)
@timer
async def retrieve_today_memory(self):
if not self.today_obs_top_k:
return
if not self.chat_messages:
self.logger.warning("chat_messages is empty!")
return
message: Message = self.chat_messages[-1]
dt_handler = DatetimeHandler(message.time_created)
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],
"dt": dt_handler.datetime_format(),
}
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=message.content,
top_k=self.today_obs_top_k,
filter_dict=filter_dict)
self.set_memories(TODAY_NODES, nodes)
async def _run(self):
mock_query = "-"
self.submit_async_task(self.retrieve_not_reflected_memory, query=mock_query)
self.submit_async_task(self.retrieve_not_updated_memory, query=mock_query)
self.submit_async_task(self.retrieve_insight_memory, query=mock_query)
self.submit_async_task(self.retrieve_today_memory)
self.gather_async_result()

View file

@ -1,5 +1,3 @@
from typing import List
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
@ -12,9 +10,11 @@ class StoreMemoryWorker(MemoryBaseWorker):
def _run(self):
store_key: str = self.store_key
if self.has_content(store_key):
memory_nodes: List[MemoryNode] = self.get_context(store_key)
self.memory_store.update_memories(memory_nodes)
if store_key == "all":
self.save_memories()
elif self.has_content(store_key):
self.save_memories(store_key)
elif store_key in self.chat_kwargs:
query = self.chat_kwargs[store_key]