[test] add test for llm models

This commit is contained in:
jinli.yl 2024-06-24 15:50:29 +08:00
parent 627a93c10b
commit 31316a9080
4 changed files with 29 additions and 34 deletions

View file

@ -3,6 +3,7 @@ import inspect
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.utils.logger import Logger
@ -10,6 +11,7 @@ from memory_scope.utils.timer import Timer
class BaseModel(metaclass=ABCMeta):
model_type: ModelEnum | None = None
def __init__(self,
model_name: str,

View file

@ -1,40 +1,25 @@
from typing import List, Dict
#from llama_index.llms.dashscope import DashScope as DashScopeLLM
from llama_index.llms.dashscope import DashScope
from llama_index.core.base.llms.types import ChatMessage
from llama_index.core.base.llms.types import (
ChatResponse,
CompletionResponse,
)
from llama_index.llms.dashscope import DashScope
from memory_scope.enumeration.model_enum import ModelEnum
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.utils.timer import Timer
class BaseGenerationModel(BaseModel):
class LlamaIndexGenerationModel(BaseModel):
model_type: ModelEnum = ModelEnum.GENERATION_MODEL
MODEL_REGISTRY.batch_register([
DashScope,
])
def before_call(self, **kwargs) -> None:
pass
def after_call(self, model_response: ModelResponse | ModelResponseGen,
**kwargs) -> ModelResponse | ModelResponseGen:
pass
def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen:
pass
async def _async_call(self, **kwargs) -> ModelResponse:
pass
class LLILLM(BaseGenerationModel):
def before_call(self, **kwargs) -> None:
prompt: str = kwargs.pop("prompt", "")
messages: List[Dict[str, str]] = kwargs.pop("messages", [])
@ -91,3 +76,5 @@ class LLILLM(BaseGenerationModel):
results.details = e
return results
async def _async_call(self, **kwargs) -> ModelResponse:
pass

View file

@ -2,21 +2,25 @@ from typing import Generator, List, Dict, Any
from pydantic import BaseModel, Field
from memory_scope.enumeration.model_enum import ModelEnum
class ModelResponse(BaseModel):
text: str = Field("", description="")
embedding_results: List[List[float]] = Field([], description="")
embedding_results: List[List[float]] | List[float] = Field([], description="embedding result")
#rank_scores: Dict[int, float] = Field({}, description="The rank scores of each documents.")
rank_scores: List[Dict[str, Any]] = Field({}, description="The rank scores of each documents.")
# [{"document": "xxx", "score": 0.5}, {{"document": "yyy", "score": 0.3}}]
model_type: str = Field("", description="One of LLM, EMB, RANK.")
rank_scores: List[Dict[int, float]] = Field([], description="The rank scores of each documents. "
"key: index, value: rank score")
model_type: ModelEnum = Field("", description="One of LLM, EMB, RANK.")
status: bool = Field(True, description="Indicates whether the model call was successful.")
details: str = Field("", description=("The details information for model call, \
usually for storage of raw response or failure messages."))
details: str = Field("", description="The details information for model call, "
"usually for storage of raw response or failure messages.")
raw: Any = Field("", description="raw response from model call")
raw: Any = Field("", description=("raw response from model call"))
ModelResponseGen = Generator[ModelResponse, None, None]

View file

@ -1,5 +1,7 @@
import unittest
from memory_scope.models.base_generation_model import LLILLM
from memory_scope.models.llama_index_generation_model import LlamaIndexGenerationModel
class TestLLILLM(unittest.TestCase):
"""Tests for LLIEmbedding"""
@ -8,9 +10,9 @@ class TestLLILLM(unittest.TestCase):
config = {
"method_type": "DashScope",
"model_name": "qwen-max",
"clazz": "models.base_generation_model"
"clazz": "models.llama_index_generation_model"
}
self.llm = LLILLM(**config)
self.llm = LlamaIndexGenerationModel(**config)
def test_llm_prompt(self):
prompt = "你是谁?"
@ -21,10 +23,10 @@ class TestLLILLM(unittest.TestCase):
print(ans.text)
def test_llm_messages(self):
messages = [{"role": "system", "content": "you are a helpful assistant."},
messages = [{"role": "system", "content": "you are a helpful assistant."},
{"role": "user", "content": "你是谁?"}]
ans = self.llm.call(
stream=False,
messages=messages
)
print(ans.text)
print(ans.text)