[dev] format model response

This commit is contained in:
jinli.yl 2024-06-26 18:27:39 +08:00
parent 7c30c104fd
commit 02bb9718be
6 changed files with 23 additions and 43 deletions

View file

@ -1,3 +1 @@
from memory_scope.utils.registry import Registry
MODEL_REGISTRY = Registry("models")

View file

@ -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)

View file

@ -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))

View file

@ -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", "")

View file

@ -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

View file

@ -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__