delete default service params

This commit is contained in:
jinli.yl 2024-07-11 21:09:23 +08:00
parent 8282350ab4
commit f0a13ae801
19 changed files with 141 additions and 120 deletions

View file

@ -11,7 +11,6 @@ memory_chat:
memory_service:
memory_chat_service:
class: memory.service.chat_memory_service
history_msg_count: 32
contextual_msg_count: 6
memory_operations:
read_message:
@ -30,11 +29,11 @@ memory_service:
workflow: info_filter,load_memory1,[get_observation|get_observation_with_time],contra_repeat,store_memory
description: "write observation memories of the user"
interval_time: 5
# summary_memory:
# class: memory.operation.summary_memory
# workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory
# description: "summary observation memories of the user"
# interval_time: 60
summary_memory:
class: memory.operation.summary_memory
workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory
description: "summary observation memories of the user"
interval_time: 30
worker:
dummy:
@ -76,6 +75,7 @@ worker:
info_filter:
class: memory.worker.write.info_filter_worker
generation_model: dashscope_generation
preserved_scores: 2,3
info_filter_msg_max_size: 200
generation_model_top_k: 1
load_memory1:
@ -110,7 +110,7 @@ worker:
get_reflection_subject:
class: memory.worker.summary.get_reflection_subject_worker
retrieve_top_k: 100
reflect_obs_cnt_threshold: 32
reflect_obs_cnt_threshold: 10
generation_model_top_k: 1
update_insight:
class: memory.worker.summary.update_insight_worker

View file

@ -1,5 +1,7 @@
from abc import ABCMeta, abstractmethod
from memory_scope.memory.service.base_memory_service import BaseMemoryService
class BaseMemoryChat(metaclass=ABCMeta):
@ -10,5 +12,9 @@ class BaseMemoryChat(metaclass=ABCMeta):
:return:
"""
@property
def memory_service(self) -> BaseMemoryService:
raise NotImplementedError
def run(self):
pass

View file

@ -76,7 +76,7 @@ class CliMemoryChat(BaseMemoryChat):
self._generation_model = G_CONTEXT.model_dict[self._generation_model]
return self._generation_model
def chat_with_memory(self, query: str, remember_response:bool=False) -> ModelResponse | ModelResponseGen:
def chat_with_memory(self, query: str, remember_response: bool = False) -> ModelResponse | ModelResponseGen:
new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=self.human_name, content=query)
self.memory_service.add_messages(new_message)
@ -107,7 +107,7 @@ class CliMemoryChat(BaseMemoryChat):
assert not self.stream
generated.message.role_name = self.assistant_name
self.memory_service.add_messages(generated.message)
# return response or generator
return generated
@ -217,5 +217,3 @@ class CliMemoryChat(BaseMemoryChat):
traceback.print_exc()
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

@ -4,8 +4,10 @@ system_prompt:
en: |
You are a helpful assistant, your name is MemoryScope.
# 请记住以下信息,他们可以帮助更好地理解用户的问题。
memory_prompt:
cn: |
请记住以下信息,他们可以帮助更好地理解用户的问题。
如果用户问题和以下信息没有关联,请忘记这些信息;如果用户问题和以下信息有关联,请记住这些信息,他们可以帮助更好地理解用户的问题。
en: |
Please remember the following information, as they can help better understand the user's question.
If the user's question is not related to the following information, please disregard it; if the user's question is related to the following information, please retain it, as it can help better understand the user's query.

View file

@ -1,6 +1,8 @@
import datetime
import sys
import questionary
sys.path.append(".") # noqa: E402
import json
@ -16,7 +18,7 @@ from memory_scope.enumeration.model_enum import ModelEnum
from memory_scope.utils.global_context import G_CONTEXT
from memory_scope.utils.logger import Logger
from memory_scope.utils.timer import timer
from memory_scope.utils.tool_functions import init_instance_by_config
from memory_scope.utils.tool_functions import init_instance_by_config, camelcase_to_underscore
class MemoryScope(object):
@ -24,7 +26,8 @@ class MemoryScope(object):
def __init__(self):
self.config: Dict[str, Any] = {}
datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S')
self.logger: Logger = Logger.get_logger(f"cli_job_{datetime_suffix}", to_stream=False)
class_name = camelcase_to_underscore(self.__class__.__name__)
self.logger: Logger = Logger.get_logger(f"{class_name}_{datetime_suffix}", to_stream=False)
def load_config(self, path: str):
with open(path) as f:
@ -40,7 +43,7 @@ class MemoryScope(object):
@staticmethod
def shutdown():
print('Gracefully executing the shutdown function...')
questionary.print('Gracefully executing the shutdown function...')
G_CONTEXT.memory_store.close()
G_CONTEXT.monitor.close()
G_CONTEXT.thread_pool.shutdown()
@ -80,24 +83,21 @@ class MemoryScope(object):
# set worker config
G_CONTEXT.worker_config = self.config["worker"]
@property
def default_service(self):
return list(G_CONTEXT.memory_service_dict.values())[0]
@property
def default_chat_handle(self):
return list(G_CONTEXT.memory_chat_dict.values())[0]
@property
def default_service(self):
return self.default_chat_handle.memory_service
class CliJob(MemoryScope):
def run(self, config: str):
self.load_config(config)
self.init_global_content_by_config()
# with G_CONTEXT.thread_pool:
memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0]
memory_chat.run()
self.default_chat_handle.run()
if __name__ == "__main__":

View file

@ -159,3 +159,8 @@ DATATIME_KEY_MAP = {
"Weekday": "weekday",
}
}
TIME_INFER_WORD = {
LanguageEnum.CN: "推断时间",
LanguageEnum.EN: "Inference time"
}

View file

@ -6,7 +6,7 @@ from memory_scope.utils.tool_functions import init_instance_by_config
class ChatMemoryService(BaseMemoryService):
def __init__(self, history_msg_count: int = 32, contextual_msg_count: int = 6, **kwargs):
def __init__(self, history_msg_count: int = 100, contextual_msg_count: int = 6, **kwargs):
super().__init__(**kwargs)
self.history_msg_count: int = history_msg_count
self.contextual_msg_count: int = contextual_msg_count

View file

@ -76,23 +76,26 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
return self.get_context(CONTEXT_MEMORY_DICT)
def get_memories(self, keys: str | List[str]) -> List[MemoryNode]:
memories: List[MemoryNode] = []
memories: Dict[str, MemoryNode] = {}
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
memories.update({x: self.contex_memory_dict[x] for x in memory_ids})
return list(memories.values())
def set_memories(self, key: str, nodes: MemoryNode | List[MemoryNode]):
def set_memories(self, key: str, nodes: MemoryNode | List[MemoryNode], log_repeat: bool = True):
if nodes is None:
nodes = []
elif isinstance(nodes, MemoryNode):
nodes = [nodes]
for node in nodes:
if node.memory_id in self.contex_memory_dict:
if log_repeat:
self.logger.warning(f"repeated_id memory id={node.memory_id} content={node.content} "
f"status={node.status}")
continue
self.contex_memory_dict[node.memory_id] = node
self.logger.info(f"add to memory context memory id={node.memory_id} content={node.content} "

View file

@ -20,23 +20,26 @@ class PrintMemoryWorker(MemoryBaseWorker):
i = 0
j = 0
k = 0
expired_content_set = set()
for node in memory_node_list:
dt_handler = DatetimeHandler(node.timestamp)
dt = dt_handler.datetime_format("%Y%m%d %H:%M:%S")
line = f"{dt} {node.content}"
if MemoryNodeStatus(node.status) is MemoryNodeStatus.EXPIRED:
i += 1
line = f" {i} {node.content}"
expired_content_list.append(line)
if node.content in expired_content_set:
continue
else:
expired_content_set.add(node.content)
i += 1
expired_content_list.append(f" {i} {line}")
elif MemoryTypeEnum(node.memory_type) in [MemoryTypeEnum.OBSERVATION, MemoryTypeEnum.OBS_CUSTOMIZED]:
j += 1
dt_handler = DatetimeHandler(node.timestamp)
dt = dt_handler.datetime_format("%Y%m%d %H:%M:%S")
line = f" {j} {dt} {node.content}"
obs_content_list.append(line)
obs_content_list.append(f" {j} {line}")
elif MemoryTypeEnum(node.memory_type) is MemoryTypeEnum.INSIGHT:
k += 1
line = f" {k} {node.content}"
insight_content_list.append(line)
insight_content_list.append(f" {k} {line}")
obs_content = "\n".join(obs_content_list)
insight_content = "\n".join(insight_content_list)

View file

@ -33,4 +33,4 @@ class SemanticRankWorker(MemoryBaseWorker):
memory_node_list = sorted(memory_node_list, key=lambda n: n.score_rank, reverse=True)
for node in memory_node_list:
self.logger.info(f"rank_stage: content={node.content} score={node.score_rank}")
self.set_memories(RANKED_MEMORY_NODES, memory_node_list)
self.set_memories(RANKED_MEMORY_NODES, memory_node_list, log_repeat=False)

View file

@ -33,6 +33,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
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.")
self.continue_run = False
return
# get profile_keys

View file

@ -83,10 +83,8 @@ class ContraRepeatWorker(MemoryBaseWorker):
node: MemoryNode = all_obs_nodes[idx]
if keep_flag != self.get_language_value(NONE_WORD):
node.status = MemoryNodeStatus.EXPIRED.value
self.logger.info(f"contra_repeat stage: {node.content} {node.status}")
merge_obs_nodes.append(node)
# forbid keyword
self.logger.info(f"contra_repeat stage: {node.content} {node.status}")
# save context
self.set_memories(MERGE_OBS_NODES, merge_obs_nodes)
self.set_memories(MERGE_OBS_NODES, merge_obs_nodes, log_repeat=False)

View file

@ -1,7 +1,7 @@
from typing import List
from memory_scope.constants.common_constants import NEW_OBS_NODES, TIME_INFER
from memory_scope.constants.language_constants import REPEATED_WORD, NONE_WORD, COLON_WORD
from memory_scope.constants.language_constants import REPEATED_WORD, NONE_WORD, COLON_WORD, TIME_INFER_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
@ -30,6 +30,8 @@ class GetObservationWorker(MemoryBaseWorker):
if time_infer:
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()})
obs_content = (f"{obs_content} ({self.get_language_value(TIME_INFER_WORD)}"
f"{self.get_language_value(COLON_WORD)} {time_infer})")
return MemoryNode(user_name=self.user_name,
target_name=self.target_name,

View file

@ -69,7 +69,7 @@ class InfoFilterWorker(MemoryBaseWorker):
continue
score = info_score[0]
if score in ("3",):
if score in self.preserved_scores:
msg.meta_data["info_score"] = score
filtered_messages.append(msg)
self.chat_messages = filtered_messages

View file

@ -6,7 +6,7 @@ from pydantic import Field, BaseModel
class MemoryNode(BaseModel):
memory_id: str = Field(str(uuid4()), description="unique id for memory")
memory_id: str = Field(default_factory=lambda: uuid4().hex, description="unique id for memory")
user_name: str = Field("", description="the user who owns the memory")
@ -33,7 +33,8 @@ class MemoryNode(BaseModel):
vector: List[float] = Field([], description="content embedding result, return empty")
timestamp: int = Field(int(datetime.datetime.now().timestamp()), description="timestamp of the memory node")
timestamp: int = Field(default_factory=lambda: int(datetime.datetime.now().timestamp()),
description="timestamp of the memory node")
dt: str = Field("", description="dt of the memory node")

View file

@ -1,3 +1,4 @@
import warnings
from typing import Dict, List, Any, Optional, cast
from llama_index.core import VectorStoreIndex
@ -137,10 +138,13 @@ class LlamaIndexEsMemoryStoreSync(BaseMemoryStore):
es_url=es_url,
retrieval_strategy=_AsyncDenseVectorStrategy(hybrid=use_hybrid),
**kwargs)
# use /dev/null
with open(os.devnull, 'w') as devnull:
# TODO The llamaIndex utilizes some deprecated functions, hence langchain logs warning messages. By
# adding the following lines of code, the display of deprecated information is suppressed.
with warnings.catch_warnings():
warnings.simplefilter("ignore")
self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store,
embed_model=self.embedding_model.model)
self.index.build_index_from_nodes([TextNode(text="text")])
self.logger = Logger.get_logger()

View file

@ -1,13 +1,19 @@
"""Elasticsearch vector store."""
import asyncio
from logging import getLogger
from typing import Any, Callable, Dict, List, Literal, Optional, Union
import nest_asyncio
import numpy as np
from elasticsearch import AsyncElasticsearch, Elasticsearch
from elasticsearch.helpers.vectorstore import (
AsyncBM25Strategy,
AsyncSparseVectorStrategy,
AsyncDenseVectorStrategy,
AsyncRetrievalStrategy,
DistanceMetric,
)
from elasticsearch.helpers.vectorstore import VectorStore
from llama_index.core.bridge.pydantic import PrivateAttr
from llama_index.core.schema import BaseNode, MetadataMode, TextNode
from llama_index.core.vector_stores.types import (
@ -21,20 +27,10 @@ from llama_index.core.vector_stores.utils import (
metadata_dict_to_node,
node_to_metadata_dict,
)
from elasticsearch.helpers.vectorstore import AsyncVectorStore, VectorStore
from elasticsearch.helpers.vectorstore import (
AsyncBM25Strategy,
AsyncSparseVectorStrategy,
AsyncDenseVectorStrategy,
AsyncRetrievalStrategy,
DistanceMetric,
)
from llama_index.vector_stores.elasticsearch.utils import (
get_user_agent,
)
logger = getLogger(__name__)
DISTANCE_STRATEGIES = Literal[
@ -43,13 +39,14 @@ DISTANCE_STRATEGIES = Literal[
"EUCLIDEAN_DISTANCE",
]
def get_elasticsearch_client(
url: Optional[str] = None,
cloud_id: Optional[str] = None,
api_key: Optional[str] = None,
username: Optional[str] = None,
password: Optional[str] = None,
use_async: Optional[bool] = False,
url: Optional[str] = None,
cloud_id: Optional[str] = None,
api_key: Optional[str] = None,
username: Optional[str] = None,
password: Optional[str] = None,
use_async: Optional[bool] = False,
) -> AsyncElasticsearch:
if url and cloud_id:
raise ValueError(
@ -126,7 +123,7 @@ def _to_llama_similarities(scores: List[float]) -> List[float]:
def _mode_must_match_retrieval_strategy(
mode: VectorStoreQueryMode, retrieval_strategy: AsyncRetrievalStrategy
mode: VectorStoreQueryMode, retrieval_strategy: AsyncRetrievalStrategy
) -> None:
"""
Different retrieval strategies require different ways of indexing that must be known at the
@ -241,20 +238,21 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
retrieval_strategy: AsyncRetrievalStrategy
_store = PrivateAttr()
def __init__(
self,
index_name: str,
es_client: Optional[Any] = None,
es_url: Optional[str] = None,
es_cloud_id: Optional[str] = None,
es_api_key: Optional[str] = None,
es_user: Optional[str] = None,
es_password: Optional[str] = None,
text_field: str = "content",
vector_field: str = "embedding",
batch_size: int = 200,
distance_strategy: Optional[DISTANCE_STRATEGIES] = "COSINE",
retrieval_strategy: Optional[AsyncRetrievalStrategy] = None,
self,
index_name: str,
es_client: Optional[Any] = None,
es_url: Optional[str] = None,
es_cloud_id: Optional[str] = None,
es_api_key: Optional[str] = None,
es_user: Optional[str] = None,
es_password: Optional[str] = None,
text_field: str = "content",
vector_field: str = "embedding",
batch_size: int = 200,
distance_strategy: Optional[DISTANCE_STRATEGIES] = "COSINE",
retrieval_strategy: Optional[AsyncRetrievalStrategy] = None,
) -> None:
nest_asyncio.apply()
@ -310,13 +308,13 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
def close(self) -> None:
return self._store.close()
def add(
self,
nodes: List[BaseNode],
*,
create_index_if_not_exists: bool = True,
**add_kwargs: Any,
self,
nodes: List[BaseNode],
*,
create_index_if_not_exists: bool = True,
**add_kwargs: Any,
) -> List[str]:
"""
Add nodes to Elasticsearch index.
@ -335,15 +333,15 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
ImportError: If elasticsearch['async'] python package is not installed.
BulkIndexError: If AsyncElasticsearch async_bulk indexing fails.
"""
return self.sync_add(nodes, create_index_if_not_exists=create_index_if_not_exists)
def sync_add(
self,
nodes: List[BaseNode],
*,
create_index_if_not_exists: bool = True,
**add_kwargs: Any,
self,
nodes: List[BaseNode],
*,
create_index_if_not_exists: bool = True,
**add_kwargs: Any,
) -> List[str]:
"""
Asynchronous method to add nodes to Elasticsearch index.
@ -419,19 +417,19 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
return self._store.delete(query={"term": {"_id": ref_doc_id}}, **delete_kwargs)
def query(
self,
query: VectorStoreQuery,
custom_query: Optional[
Callable[[Dict, Union[VectorStoreQuery, None]], Dict]
] = None,
es_filter: Optional[List[Dict]] = None,
**kwargs: Any,
self,
query: VectorStoreQuery,
custom_query: Optional[
Callable[[Dict, Union[VectorStoreQuery, None]], Dict]
] = None,
es_filter: Optional[List[Dict]] = None,
**kwargs: Any,
) -> VectorStoreQueryResult:
"""
Query index for top k most similar nodes.
Args:
query_embedding (List[float]): query embedding
query (List[float]): query embedding
custom_query: Optional. custom query function that takes in the es query
body and returns a modified query body.
This can be used to add additional query
@ -447,16 +445,16 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
Exception: If Elasticsearch query fails.
"""
return self.sync_query(query, custom_query, es_filter, **kwargs)
return self.sync_query(query, custom_query, es_filter, **kwargs)
def sync_query(
self,
query: VectorStoreQuery,
custom_query: Optional[
Callable[[Dict, Union[VectorStoreQuery, None]], Dict]
] = None,
es_filter: Optional[List[Dict]] = None,
**kwargs: Any,
self,
query: VectorStoreQuery,
custom_query: Optional[
Callable[[Dict, Union[VectorStoreQuery, None]], Dict]
] = None,
es_filter: Optional[List[Dict]] = None,
**kwargs: Any,
) -> VectorStoreQueryResult:
"""
Asynchronous query index for top k most similar nodes.
@ -488,7 +486,7 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
query=query.query_str,
query_vector=query.query_embedding,
k=query.similarity_top_k,
num_candidates=100, # query.similarity_top_k * 10,
num_candidates=100, # query.similarity_top_k * 10,
filter=filter,
custom_query=custom_query,
)
@ -531,14 +529,13 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
top_k_scores.append(hit.get("_rank", hit["_score"]))
if (
isinstance(self.retrieval_strategy, AsyncDenseVectorStrategy)
and self.retrieval_strategy.hybrid
isinstance(self.retrieval_strategy, AsyncDenseVectorStrategy)
and self.retrieval_strategy.hybrid
):
total_rank = sum(top_k_scores)
top_k_scores = [(total_rank - rank) / total_rank for rank in top_k_scores]
# top_k_scores = [total_rank - rank / total_rank for rank in top_k_scores]
return VectorStoreQueryResult(
nodes=top_k_nodes,
ids=top_k_ids,

View file

@ -143,9 +143,9 @@ class DatetimeHandler(object):
@classmethod
def extract_date_parts(cls, input_string: str) -> dict:
func_name = f"extract_date_parts_{G_CONTEXT.language}"
func_name = f"extract_date_parts_{G_CONTEXT.language.value}"
if not hasattr(cls, func_name):
cls.logger.warning(f"language={G_CONTEXT.language} needs to complete extract_date_parts func!")
cls.logger.warning(f"language={G_CONTEXT.language.value} needs to complete extract_date_parts func!")
return {}
return getattr(cls, func_name)(input_string=input_string)
@ -166,9 +166,9 @@ class DatetimeHandler(object):
@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}"
func_name = f"format_time_by_extract_time_{G_CONTEXT.language.value}"
if not hasattr(cls, func_name):
cls.logger.warning(f"language={G_CONTEXT.language} needs to complete format_time_by_extract_time func!")
cls.logger.warning(f"language={G_CONTEXT.language.value} needs to complete format_time_by_extract_time func!")
return ""
return getattr(cls, func_name)(extract_time_dict, meta_data)

View file

@ -13,4 +13,5 @@ pydantic~=2.7.1
dashscope~=1.19.1
elasticsearch~=8.14.0
pyyaml~=6.0.1
ray
ray~=2.31.0
numpy~=1.26.4