context passing

This commit is contained in:
青轩 2024-08-13 15:25:13 +08:00
parent eacf6b5f85
commit 22683d29ff
7 changed files with 40 additions and 9 deletions

View file

@ -1,7 +1,6 @@
from dataclasses import dataclass, field
from typing import Literal, Dict
@dataclass
class Arguments(object):
language: Literal["cn", "en"] = field(default="en", metadata={"help": "support en & cn now"})

View file

@ -1,4 +1,5 @@
from concurrent.futures import ThreadPoolExecutor
from datetime import datetime
from memoryscope.core.chat.base_memory_chat import BaseMemoryChat
from memoryscope.core.config.config_manager import ConfigManager
@ -32,6 +33,10 @@ class MemoryScope(ConfigManager):
self.logger.warning("If a semantic ranking model is not available, MemoryScope will use cosine similarity "
"scoring as a substitute. However, the ranking effectiveness will be somewhat "
"compromised.")
self.context.memory_scope_uuid = datetime.now().strftime(global_conf["logger_name_time_suffix"])
# set context_initialized
self.context.context_initialized = True
# init memory_chat
memory_chat_conf_dict = self.config["memory_chat"]

View file

@ -2,8 +2,9 @@ from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass, field
from memoryscope.enumeration.language_enum import LanguageEnum
from memoryscope.core.utils.singleton import singleton
@singleton
@dataclass
class MemoryscopeContext(object):
"""
@ -27,3 +28,14 @@ class MemoryscopeContext(object):
worker_conf_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> worker_conf"})
meta_data: dict = field(default_factory=lambda: {})
memory_scope_uuid: str = ""
context_initialized: bool = False
def get_ms_context():
ms_context = MemoryscopeContext()
if ms_context.context_initialized:
return ms_context
else:
raise RuntimeError("MemoryscopeContext is not initialized yet. Please initialize it first.")

View file

@ -8,6 +8,8 @@ from memoryscope.core.utils.registry import Registry
from memoryscope.core.utils.timer import Timer
from memoryscope.enumeration.model_enum import ModelEnum
from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen
from memoryscope.core.memoryscope_context import MemoryscopeContext
from memoryscope.core.memoryscope_context import get_ms_context
MODEL_REGISTRY = Registry("models")
@ -32,6 +34,7 @@ class BaseModel(metaclass=ABCMeta):
self.retry_interval: float = retry_interval
self.kwargs_filter: bool = kwargs_filter
self.raise_exception: bool = raise_exception
self.context: MemoryscopeContext = get_ms_context()
self.kwargs: dict = kwargs
self._model: Any = None

View file

@ -15,6 +15,12 @@ from .tool_functions import (
cosine_similarity
)
def get_context():
from memoryscope import MemoryscopeContext
return MemoryscopeContext()
__all__ = [
"DatetimeHandler",
"Logger",

View file

@ -84,7 +84,6 @@ class Logger(logging.Logger):
def format_chat_message(self, message):
buf = '\n'
buf += f"++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++\n"
buf += f"LM Input:\n"
for chat_message in message.meta_data['data']['messages']:
buf += chat_message.content
@ -93,13 +92,11 @@ class Logger(logging.Logger):
buf += f"LM Output:\n"
buf += message.message.content
buf += '\n'
buf += f"++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++\n"
buf += '\n'
return buf
return self.wrap_in_box(buf)
def format_rank_message(self, model_response):
buf = '\n'
buf += f"++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++\n"
buf += f"Query Input:\n"
buf += model_response.meta_data['data']['query_str']
buf += '\n'
@ -112,9 +109,8 @@ class Logger(logging.Logger):
node_text = node.text
buf += f"Score {score} | Rank {rank} | {node_text}\n"
buf += '\n'
buf += f"++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++\n"
buf += '\n'
return buf
return self.wrap_in_box(buf)
def _add_file_handler(self):
"""
@ -239,4 +235,5 @@ class Logger(logging.Logger):
@staticmethod
def append_timestamp(name: str) -> str:
return f"{name}_{datetime.now().strftime(r'%Y%m%d_%H%M')}"
from memoryscope.core.memoryscope_context import get_ms_context
return f"{name}_{get_ms_context().memory_scope_uuid}"

View file

@ -0,0 +1,9 @@
def singleton(cls):
_instance = {}
def _singleton(*args, **kargs):
if cls not in _instance:
_instance[cls] = cls(*args, **kargs)
return _instance[cls]
return _singleton