diff --git a/memory_scope/models/llama_index_generation_model.py b/memory_scope/models/llama_index_generation_model.py index f3b48cb5..a38c8ffd 100644 --- a/memory_scope/models/llama_index_generation_model.py +++ b/memory_scope/models/llama_index_generation_model.py @@ -69,4 +69,14 @@ class LlamaIndexGenerationModel(BaseModel): return results async def _async_call(self, **kwargs) -> ModelResponse: + assert "prompt" in self.data or "messages" in self.data + results = ModelResponse(m_type=self.m_type) + + if "prompt" in self.data: + response = await self.model.acomplete(**self.data) + else: + response = await self.model.achat(**self.data) + results.raw = response + return results + raise NotImplementedError diff --git a/memory_scope/models/llama_index_rank_model.py b/memory_scope/models/llama_index_rank_model.py index 71e4acef..e5cefa12 100644 --- a/memory_scope/models/llama_index_rank_model.py +++ b/memory_scope/models/llama_index_rank_model.py @@ -42,7 +42,7 @@ class LlamaIndexRankModel(BaseModel): return ModelResponse(m_type=self.m_type, raw=self.model.postprocess_nodes(**self.data)) async def _async_call(self, **kwargs) -> ModelResponse: - raise NotImplementedError + return self._call(**kwargs) def _get_documents_mapping(self, documents): self.documents_map = {} diff --git a/tests/models/test_models_lli_generation.py b/tests/models/test_models_lli_generation.py index 6f7b43cd..20b6a543 100644 --- a/tests/models/test_models_lli_generation.py +++ b/tests/models/test_models_lli_generation.py @@ -4,7 +4,7 @@ sys.path.append(".") # noqa: E402 import unittest import time - +import asyncio from memory_scope.scheme.message import Message from memory_scope.models.llama_index_generation_model import LlamaIndexGenerationModel from memory_scope.utils.logger import Logger @@ -17,7 +17,10 @@ class TestLLILLM(unittest.TestCase): config = { "module_name": "dashscope_generation", "model_name": "qwen-max", - "clazz": "models.llama_index_generation_model" + "clazz": "models.llama_index_generation_model", + "max_tokens": 2000, + "top_k": 1, + "seed": 1234, } self.llm = LlamaIndexGenerationModel(**config) self.logger = Logger.get_logger() @@ -29,7 +32,7 @@ class TestLLILLM(unittest.TestCase): def test_llm_messages(self): messages = [Message(role="system", content="you are a helpful assistant."), - Message(role="user", content="你是谁?")] + Message(role="user", content="你如何看待黄金上涨?")] ans = self.llm.call(stream=False, messages=messages) self.logger.info(ans.message.content) @@ -53,3 +56,11 @@ class TestLLILLM(unittest.TestCase): sys.stdout.flush() time.sleep(0.1) self.logger.info("-----end-----") + + def test_async_llm_messages(self): + + messages = [Message(role="system", content="you are a helpful assistant."), + Message(role="user", content="你如何看待黄金上涨?")] + + ans = asyncio.run(self.llm.async_call(messages=messages)) + self.logger.info(ans.message.content) \ No newline at end of file diff --git a/tests/models/test_models_lli_rank.py b/tests/models/test_models_lli_rank.py index 1c3b76a5..d3865e25 100644 --- a/tests/models/test_models_lli_rank.py +++ b/tests/models/test_models_lli_rank.py @@ -1,5 +1,5 @@ import unittest - +import asyncio from memory_scope.models.llama_index_rank_model import LlamaIndexRankModel @@ -19,7 +19,15 @@ class TestLLIReRank(unittest.TestCase): documents = ["您吃了吗?", "吃了吗您?"] result = self.reranker.call( - stream=False, documents=documents, query=query) print(result) + + def test_async_rerank(self): + query = "吃啥?" + documents = ["您吃了吗?", + "吃了吗您?"] + result = asyncio.run(self.reranker.async_call( + documents=documents, + query=query)) + print(result) \ No newline at end of file diff --git a/tests/storages/test_storages_lli_es.py b/tests/storages/test_storages_lli_es.py index 7f88bea2..34c66735 100644 --- a/tests/storages/test_storages_lli_es.py +++ b/tests/storages/test_storages_lli_es.py @@ -139,7 +139,9 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): status="invalid", memory_id="ggg567" )) - res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10) + import asyncio + res = asyncio.run(self.es_store.async_retrieve(query="hacker", filter_dict=filter_dict, top_k=10)) + #res = self.es_store.async_retrieve(query="hacker", filter_dict=filter_dict, top_k=10) print(len(res)) print(res)