mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-14 23:21:04 +00:00
[dev] modify test for llm
This commit is contained in:
parent
28624de878
commit
093946e999
3 changed files with 28 additions and 19 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)")
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue