diff --git a/memory_scope/models/__init__.py b/memory_scope/models/__init__.py index 5a9d5579..8b137891 100644 --- a/memory_scope/models/__init__.py +++ b/memory_scope/models/__init__.py @@ -1,3 +1 @@ -from memory_scope.utils.registry import Registry -MODEL_REGISTRY = Registry("models") diff --git a/memory_scope/models/base_model.py b/memory_scope/models/base_model.py index 28b12fe5..038cde81 100644 --- a/memory_scope/models/base_model.py +++ b/memory_scope/models/base_model.py @@ -3,11 +3,13 @@ import time from abc import abstractmethod, ABCMeta from memory_scope.enumeration.model_enum import ModelEnum -from memory_scope.models import MODEL_REGISTRY -from memory_scope.models.response import ModelResponse, ModelResponseGen +from memory_scope.models.model_response import ModelResponse, ModelResponseGen from memory_scope.utils.logger import Logger +from memory_scope.utils.registry import Registry from memory_scope.utils.timer import Timer +MODEL_REGISTRY = Registry("models") + class BaseModel(metaclass=ABCMeta): m_type: ModelEnum | None = None @@ -70,8 +72,8 @@ class BaseModel(metaclass=ABCMeta): :param kwargs: :return: """ - self.before_call(stream=stream, **kwargs) with Timer(self.__class__.__name__, log_time=False) as t: + self.before_call(stream=stream, **kwargs) for i in range(self.max_retries): try: model_response = self._call(stream=stream, **kwargs) @@ -97,8 +99,8 @@ class BaseModel(metaclass=ABCMeta): :param kwargs: :return: """ - self.before_call(**kwargs) 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) diff --git a/memory_scope/models/llama_index_embedding_model.py b/memory_scope/models/llama_index_embedding_model.py index a9397116..bb4d35b4 100644 --- a/memory_scope/models/llama_index_embedding_model.py +++ b/memory_scope/models/llama_index_embedding_model.py @@ -2,20 +2,15 @@ from typing import List from llama_index.embeddings.dashscope import DashScopeEmbedding -from memory_scope.models import MODEL_REGISTRY -from memory_scope.models.base_model import BaseModel -from memory_scope.models.response import ModelResponse, ModelResponseGen from memory_scope.enumeration.model_enum import ModelEnum +from memory_scope.models.base_model import BaseModel, MODEL_REGISTRY +from memory_scope.models.model_response import ModelResponse class LlamaIndexEmbeddingModel(BaseModel): m_type: ModelEnum = ModelEnum.EMBEDDING_MODEL - MODEL_REGISTRY.batch_register( - [ - DashScopeEmbedding, - ] - ) + MODEL_REGISTRY.register("dashscope_embedding", DashScopeEmbedding) def before_call(self, **kwargs): text: str | List[str] = kwargs.pop("text", "") @@ -42,16 +37,11 @@ class LlamaIndexEmbeddingModel(BaseModel): :param kwargs: :return: """ - return ModelResponse( - m_type=self.m_type, raw=self.model.get_text_embedding_batch(**self.data) - ) + return ModelResponse(m_type=self.m_type, raw=self.model.get_text_embedding_batch(**self.data)) async def _async_call(self, **kwargs) -> ModelResponse: """ :param kwargs: :return: """ - return ModelResponse( - m_type=self.m_type, - raw=await self.model.aget_text_embedding_batch(**self.data), - ) + return ModelResponse(m_type=self.m_type, raw=await self.model.aget_text_embedding_batch(**self.data)) diff --git a/memory_scope/models/llama_index_generation_model.py b/memory_scope/models/llama_index_generation_model.py index e2662c12..56f0b88b 100644 --- a/memory_scope/models/llama_index_generation_model.py +++ b/memory_scope/models/llama_index_generation_model.py @@ -1,24 +1,17 @@ from typing import List, Dict -from llama_index.core.base.llms.types import ( - ChatMessage, - ChatResponse, - CompletionResponse, -) + +from llama_index.core.base.llms.types import ChatMessage, ChatResponse, CompletionResponse from llama_index.llms.dashscope import DashScope -from enumeration.model_enum import ModelEnum -from . import MODEL_REGISTRY -from .base_model import BaseModel -from .response import ModelResponse, ModelResponseGen +from memory_scope.enumeration.model_enum import ModelEnum +from memory_scope.models.base_model import BaseModel, MODEL_REGISTRY +from memory_scope.models.model_response import ModelResponse, ModelResponseGen class LlamaIndexGenerationModel(BaseModel): m_type: ModelEnum = ModelEnum.GENERATION_MODEL - # TODO rename module name at xianzhe - MODEL_REGISTRY.batch_register([ - DashScope, - ]) + MODEL_REGISTRY.register("dashscope_generation", DashScope) def before_call(self, **kwargs) -> None: prompt: str = kwargs.pop("prompt", "") diff --git a/memory_scope/models/llama_index_rerank_model.py b/memory_scope/models/llama_index_rerank_model.py index 144a3b69..1a19e6aa 100644 --- a/memory_scope/models/llama_index_rerank_model.py +++ b/memory_scope/models/llama_index_rerank_model.py @@ -4,19 +4,15 @@ from llama_index.core.data_structs import Node from llama_index.core.schema import NodeWithScore from llama_index.postprocessor.dashscope_rerank import DashScopeRerank -from models import MODEL_REGISTRY -from models.base_model import BaseModel -from models.response import ModelResponse, ModelResponseGen -from enumeration.model_enum import ModelEnum - +from memory_scope.enumeration.model_enum import ModelEnum +from memory_scope.models.base_model import BaseModel, MODEL_REGISTRY +from memory_scope.models.model_response import ModelResponse class LlamaIndexRerankModel(BaseModel): m_type: ModelEnum = ModelEnum.RANK_MODEL - MODEL_REGISTRY.batch_register([ - DashScopeRerank - ]) + MODEL_REGISTRY.register("dashscope_rank", DashScopeRerank) def before_call(self, **kwargs) -> None: assert "query" in kwargs or "documents" in kwargs diff --git a/memory_scope/utils/registry.py b/memory_scope/utils/registry.py index 9939d3ab..88807387 100644 --- a/memory_scope/utils/registry.py +++ b/memory_scope/utils/registry.py @@ -10,7 +10,8 @@ class Registry(object): self.name: str = name self.module_dict: Dict[str, Any] = {} - def register(self, module: Any, module_name: str = None): + def register(self, module_name: str = None, module: Any = None): + assert module is not None if module_name is None: module_name = module.__name__