[dev] rename constants name

This commit is contained in:
jinli.yl 2024-06-21 11:39:37 +08:00
parent cba7eda6f2
commit a577bb4ce8
10 changed files with 14 additions and 24 deletions

View file

@ -9,7 +9,6 @@
]
},
"memory_chat": {
"memory_user_name": "用户",
"clazz": "chat.memory_chat",
"retrieve": "parse_params,load_profile,extract_time,es_similar,es_keyword,semantic_rank,fuse_rerank",
"generation_model": "dashscope_generation",

View file

@ -5,10 +5,8 @@ from memory_scope.chat.memory_service import MemoryService
class BaseMemoryChat(metaclass=ABCMeta):
def __init__(self, chat_name: str, memory_user_name: str, **kwargs):
self.memory_service = MemoryService(chat_name=chat_name,
memory_user_name=memory_user_name,
**kwargs)
def __init__(self, chat_name: str, **kwargs):
self.memory_service = MemoryService(chat_name=chat_name, **kwargs)
@abstractmethod
def chat_with_memory(self, query: str):

View file

@ -3,10 +3,10 @@ from typing import List
from memory_scope.chat.base_memory_chat import BaseMemoryChat
from memory_scope.chat.global_context import GLOBAL_CONTEXT
from memory_scope.definition.message import Message
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
from memory_scope.models.base_model import BaseModel
from memory_scope.prompts.memory_chat_prompt import SYSTEM_PROMPT, MEMORY_PROMPT
from memory_scope.scheme.message import Message
class MemoryChat(BaseMemoryChat):

View file

@ -1,6 +1,6 @@
from memory_scope.constants.common_constants import RELATED_MEMORIES
from memory_scope.definition.message import Message
from memory_scope.enumeration.memory_method_enum import MemoryMethodEnum
from memory_scope.scheme.message import Message
from memory_scope.utils.pipeline import Pipeline
@ -8,7 +8,6 @@ class MemoryService(object):
def __init__(self,
chat_name: str,
memory_user_name: str,
retrieve_pipeline: str,
retrieve_all_pipeline: str,
summary_short_pipeline: str,
@ -19,24 +18,20 @@ class MemoryService(object):
summary_long_minimum_count: int = 5 * 5,
**kwargs):
self.retrieve_pipeline = Pipeline(chat_name=chat_name,
user_name=memory_user_name,
memory_method_type=MemoryMethodEnum.RETRIEVE,
pipeline_str=retrieve_pipeline)
self.retrieve_all_pipeline = Pipeline(chat_name=chat_name,
user_name=memory_user_name,
memory_method_type=MemoryMethodEnum.RETRIEVE_ALL,
pipeline_str=retrieve_all_pipeline)
self.summary_short_pipeline = Pipeline(chat_name=chat_name,
user_name=memory_user_name,
memory_method_type=MemoryMethodEnum.SUMMARY_SHORT,
pipeline_str=summary_short_pipeline,
loop_interval_time=summary_short_interval_time,
loop_minimum_count=summary_short_minimum_count)
self.summary_long_pipeline = Pipeline(chat_name=chat_name,
user_name=memory_user_name,
memory_method_type=MemoryMethodEnum.SUMMARY_LONG,
pipeline_str=summary_long_pipeline,
loop_interval_time=summary_long_interval_time,

View file

@ -1,5 +1,5 @@
USER_NAME = "user_name"
RELATED_MEMORIES = "related_memories"
MESSAGES = "messages"
CHAT_NAME = "chat_name"

View file

@ -6,9 +6,9 @@ from itertools import zip_longest
from typing import Dict, Any, List
from memory_scope.chat.global_context import GLOBAL_CONTEXT
from memory_scope.constants.common_constants import MESSAGES, USER_NAME
from memory_scope.definition.message import Message
from memory_scope.constants.common_constants import MESSAGES, CHAT_NAME
from memory_scope.enumeration.memory_method_enum import MemoryMethodEnum
from memory_scope.scheme.message import Message
from memory_scope.utils.logger import Logger
from memory_scope.utils.timer import Timer
from memory_scope.worker.base_worker import BaseWorker
@ -17,7 +17,6 @@ from memory_scope.worker.base_worker import BaseWorker
class Pipeline(object):
def __init__(self,
chat_name: str,
user_name: str,
memory_method_type: MemoryMethodEnum,
pipeline_str: str,
history_msg_count: int = 3,
@ -25,7 +24,6 @@ class Pipeline(object):
loop_minimum_count: int = 20):
self.chat_name: str = chat_name
self.user_name: str = user_name
self.memory_method_type: MemoryMethodEnum = memory_method_type
self.pipeline_str: str = pipeline_str
self.history_msg_count: int = history_msg_count
@ -126,9 +124,9 @@ class Pipeline(object):
def _run(self):
self._visit_and_inject_workers()
with Timer(f"pipeline_{self.user_name}_{self.memory_method_type.value}"):
with Timer(f"pipeline_{self.chat_name}_{self.memory_method_type.value}"):
self.context[MESSAGES] = self.history_message_list + self.current_message_list
self.context[USER_NAME] = self.user_name
self.context[CHAT_NAME] = self.chat_name
for pipeline_part in self.pipeline_list:
if len(pipeline_part) == 1:

View file

@ -1,9 +1,9 @@
from typing import List
from memory_scope.chat.global_context import GLOBAL_CONTEXT
from memory_scope.constants.common_constants import MESSAGES, USER_NAME
from memory_scope.definition.message import Message
from memory_scope.constants.common_constants import MESSAGES, CHAT_NAME
from memory_scope.models.base_model import BaseModel
from memory_scope.scheme.message import Message
from memory_scope.storage.base_monitor import BaseMonitor
from memory_scope.storage.base_vector_store import BaseVectorStore
from memory_scope.worker.base_worker import BaseWorker
@ -36,8 +36,8 @@ class MemoryBaseWorker(BaseWorker):
self.set_context(MESSAGES, value)
@property
def user_name(self):
return self.get_context(USER_NAME)
def chat_name(self):
return self.get_context(CHAT_NAME)
@property
def embedding_model(self):