add async support for rank/generation/embedding models

This commit is contained in:
xianzhe.xxz 2024-07-01 10:46:54 +08:00
parent 9ce6c8ab43
commit 7f84987762
5 changed files with 38 additions and 7 deletions

View file

@ -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

View file

@ -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 = {}

View file

@ -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)

View file

@ -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)

View file

@ -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)