[dev] modify test for llm

This commit is contained in:
jinli.yl 2024-06-27 14:46:41 +08:00
parent 28624de878
commit 093946e999
3 changed files with 28 additions and 19 deletions

View file

@ -1,11 +1,14 @@
import datetime
from typing import List, Dict
from llama_index.core.base.llms.types import ChatMessage, ChatResponse, CompletionResponse
from llama_index.llms.dashscope import DashScope
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
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
from memory_scope.scheme.message import Message
class LlamaIndexGenerationModel(BaseModel):
@ -15,16 +18,16 @@ class LlamaIndexGenerationModel(BaseModel):
def before_call(self, **kwargs) -> None:
prompt: str = kwargs.pop("prompt", "")
messages: List[Dict[str, str]] = kwargs.pop("messages", [])
messages: List[Message] | List[Dict[str, str]] = kwargs.pop("messages", [])
if prompt:
input_text = prompt
input_type = 'prompt'
input_type = "prompt"
llama_input = input_text
elif messages:
input_text = messages
input_type = 'messages'
llama_input = [ChatMessage(role=x.role, content=x.content) for x in input_text]
input_type = "messages"
llama_input = [ChatMessage(role=x["role"], content=x["content"]) for x in input_text]
else:
raise RuntimeError("prompt and messages is both empty!")
@ -34,25 +37,27 @@ class LlamaIndexGenerationModel(BaseModel):
model_response: ModelResponse,
stream: bool = False,
**kwargs) -> ModelResponse | ModelResponseGen:
now_ts = datetime.datetime.now()
model_response.message = Message(role=MessageRoleEnum.ASSISTANT,
content="",
time_created=int(now_ts.timestamp()))
call_result = model_response.raw
if stream:
def gen() -> ModelResponseGen:
text = ""
for response in call_result:
delta = response.delta
text += delta
model_response.text = text
model_response.delta = delta
model_response.message.content += response.delta
model_response.delta = response.delta
yield model_response
return gen()
else:
if isinstance(call_result, CompletionResponse):
content = call_result.text
model_response.message.content = call_result.text
elif isinstance(call_result, ChatResponse):
content = call_result.message.content
model_response.message.content = call_result.message.content
else:
raise NotImplementedError
model_response.text = content
return model_response
def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen:

View file

@ -4,10 +4,11 @@ from typing import Generator, List, Dict, Any
from pydantic import BaseModel, Field
from memory_scope.enumeration.model_enum import ModelEnum
from memory_scope.scheme.message import Message
class ModelResponse(BaseModel):
text: str = Field("", description="generation model result")
message: Message | None = Field(None, description="generation model result")
delta: str = Field("", description="New text that just streamed in (only used when streaming)")

View file

@ -8,19 +8,21 @@ class TestLLILLM(unittest.TestCase):
def setUp(self):
config = {
"method_type": "DashScope",
"module_name": "dashscope_generation",
"model_name": "qwen-max",
"clazz": "models.llama_index_generation_model"
}
self.llm = LlamaIndexGenerationModel(**config)
@unittest.skip("tmp")
def test_llm_prompt(self):
prompt = "你是谁?"
ans = self.llm.call(
stream=False,
prompt=prompt
)
print(ans.text)
print(ans.message.content)
@unittest.skip("tmp")
def test_llm_messages(self):
messages = [{"role": "system", "content": "you are a helpful assistant."},
@ -29,7 +31,8 @@ class TestLLILLM(unittest.TestCase):
stream=False,
messages=messages
)
print(ans.text)
print(ans.message.content)
@unittest.skip("tmp")
def test_llm_prompt_stream(self):
prompt = "你如何看待黄金上涨?"
@ -43,8 +46,8 @@ class TestLLILLM(unittest.TestCase):
sys.stdout.write(a.delta)
sys.stdout.flush()
time.sleep(0.1)
@unittest.skip("tmp")
def test_llm_messages(self):
def test_llm_messages_stream(self):
messages = [{"role": "system", "content": "you are a helpful assistant."},
{"role": "user", "content": "你如何看待黄金上涨?"}]
ans = self.llm.call(
@ -56,4 +59,4 @@ class TestLLILLM(unittest.TestCase):
for a in ans:
sys.stdout.write(a.delta)
sys.stdout.flush()
time.sleep(0.1)
time.sleep(0.1)