mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-10 03:30:56 +00:00
add async support for rank/generation/embedding models
This commit is contained in:
parent
9ce6c8ab43
commit
7f84987762
5 changed files with 38 additions and 7 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 = {}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue