diff --git a/memory_scope/models/llama_index_generation_model.py b/memory_scope/models/llama_index_generation_model.py index 56f0b88b..af1cb59a 100644 --- a/memory_scope/models/llama_index_generation_model.py +++ b/memory_scope/models/llama_index_generation_model.py @@ -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: diff --git a/memory_scope/models/model_response.py b/memory_scope/models/model_response.py index 958356db..33047529 100644 --- a/memory_scope/models/model_response.py +++ b/memory_scope/models/model_response.py @@ -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)") diff --git a/tests/models/test_models_lli_generation.py b/tests/models/test_models_lli_generation.py index 5b3329b2..cd6cec73 100644 --- a/tests/models/test_models_lli_generation.py +++ b/tests/models/test_models_lli_generation.py @@ -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) \ No newline at end of file + time.sleep(0.1)