mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-05 08:06:15 +00:00
[dev] format model response
This commit is contained in:
parent
7c30c104fd
commit
02bb9718be
6 changed files with 23 additions and 43 deletions
|
|
@ -1,3 +1 @@
|
|||
from memory_scope.utils.registry import Registry
|
||||
|
||||
MODEL_REGISTRY = Registry("models")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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", "")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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__
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue