[dev] delete asyncio from base model

This commit is contained in:
jinli.yl 2024-06-25 12:23:39 +08:00
parent 160652b600
commit 322f48117f
8 changed files with 132 additions and 112 deletions

View file

@ -1,4 +1,3 @@
import asyncio
import inspect
import time
from abc import abstractmethod, ABCMeta
@ -11,7 +10,7 @@ from memory_scope.utils.timer import Timer
class BaseModel(metaclass=ABCMeta):
model_type: ModelEnum | None = None
m_type: ModelEnum | None = None
def __init__(self,
model_name: str,
@ -74,8 +73,12 @@ class BaseModel(metaclass=ABCMeta):
self.before_call(stream=stream, **kwargs)
with Timer(self.__class__.__name__, log_time=False) as t:
for i in range(self.max_retries):
model_response = self._call(stream=stream, **kwargs)
if not model_response.status and not stream:
try:
model_response = self._call(stream=stream, **kwargs)
except Exception as e:
model_response = ModelResponse(m_type=self.m_type, status=False, details=e.args)
if isinstance(model_response, ModelResponse) and not model_response.status:
self.logger.warning(f"call model={self.model_name} failed! cost={t.cost_str} retry_cnt={i} "
f"details={model_response.details}", stacklevel=2)
time.sleep(i * self.retry_interval)
@ -97,10 +100,14 @@ class BaseModel(metaclass=ABCMeta):
self.before_call(**kwargs)
with Timer(self.__class__.__name__, log_time=False) as t:
for i in range(self.max_retries):
model_response = await self._async_call(**kwargs)
try:
model_response = await self._async_call(**kwargs)
except Exception as e:
model_response = ModelResponse(status=False, details=e.args)
if not model_response.status:
self.logger.warning(f"async_call model={self.model_name} failed! cost={t.cost_str} retry_cnt={i} "
f"details={model_response.details}", stacklevel=2)
await asyncio.sleep(i * self.retry_interval)
time.sleep(i * self.retry_interval)
else:
return self.after_call(model_response, **kwargs)
return self.after_call(model_response=model_response, **kwargs)

View file

@ -1,13 +1,15 @@
from typing import List, Dict
from typing import List
from llama_index.embeddings.dashscope import DashScopeEmbedding
from memory_scope.enumeration.model_enum import ModelEnum
from memory_scope.models import MODEL_REGISTRY
from memory_scope.models.base_model import BaseModel
from memory_scope.models.response import ModelResponse, ModelResponseGen
from memory_scope.enumeration.model_enum import ModelEnum
from memory_scope.models.response import ModelResponse
class LlamaIndexEmbeddingModel(BaseModel):
model_type: ModelEnum = ModelEnum.EMBEDDING_MODEL
m_type: ModelEnum = ModelEnum.EMBEDDING_MODEL
MODEL_REGISTRY.batch_register([
DashScopeEmbedding,
@ -21,35 +23,28 @@ class LlamaIndexEmbeddingModel(BaseModel):
def after_call(self, model_response: ModelResponse, **kwargs) -> ModelResponse:
embeddings = model_response.raw
model_response.embedding_results = embeddings
if not embeddings:
model_response.details = "empty embeddings"
model_response.status = False
return model_response
if len(embeddings) == 1:
model_response.embedding_results = embeddings[0]
# return list[float]
embeddings = embeddings[0]
model_response.embedding_results = embeddings
return model_response
def _call(self, **kwargs) -> ModelResponse:
results = ModelResponse(model_type=self.model_type)
try:
response = self.model.get_text_embedding_batch(**self.data)
results.raw = response
results.status = True
except Exception as e:
results.details = e
results.status = False
return results
"""
:param kwargs:
:return:
"""
return ModelResponse(m_type=self.m_type, raw=self.model.get_text_embedding_batch(**self.data))
async def _async_call(self, **kwargs) -> ModelResponse:
"""
:param kwargs:
:return:
"""
results = ModelResponse(model_type=self.model_type)
try:
response = await self.model.aget_text_embedding_batch(**self.data)
results.raw = response
results.status = True
except Exception as e:
results.details = e
results.status = False
return results
return ModelResponse(m_type=self.m_type, raw=await self.model.aget_text_embedding_batch(**self.data))

View file

@ -14,8 +14,9 @@ from memory_scope.models.response import ModelResponse, ModelResponseGen
class LlamaIndexGenerationModel(BaseModel):
model_type: ModelEnum = ModelEnum.GENERATION_MODEL
m_type: ModelEnum = ModelEnum.GENERATION_MODEL
# TODO rename module name at xianzhe
MODEL_REGISTRY.batch_register([
DashScope,
])
@ -31,26 +32,24 @@ class LlamaIndexGenerationModel(BaseModel):
elif messages:
input_text = messages
input_type = 'messages'
llama_input = [ChatMessage(
role=x['role'], content=x['content']
) for x in input_text]
llama_input = [ChatMessage(role=x['role'], content=x['content']) for x in input_text]
else:
raise RuntimeError("prompt and messages is both empty!")
self.data = {
input_type: llama_input,
}
self.data = {input_type: llama_input}
def after_call(
self, model_response: ModelResponse, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen:
def after_call(self,
model_response: ModelResponse,
stream: bool = False,
**kwargs) -> ModelResponse | ModelResponseGen:
call_result = model_response.raw
if stream:
def gen() -> ModelResponseGen:
content = ""
text = ""
for response in call_result:
delta = response.delta
content += delta
model_response.text = content
text += delta
model_response.text = text
model_response.delta = delta
yield model_response
return gen()
@ -67,23 +66,19 @@ class LlamaIndexGenerationModel(BaseModel):
def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen:
assert "prompt" in self.data or "messages" in self.data
results = ModelResponse(model_type=self.model_type)
try:
if 'prompt' in self.data:
if stream:
response = self.model.stream_complete(**self.data)
else:
response = self.model.complete(**self.data)
results = ModelResponse(m_type=self.m_type)
if 'prompt' in self.data:
if stream:
response = self.model.stream_complete(**self.data)
else:
if stream:
response = self.model.stream_chat(**self.data)
else:
response = self.model.chat(**self.data)
results.raw = response
results.status = True
except Exception as e:
results.status = False
results.details = e
response = self.model.complete(**self.data)
else:
if stream:
response = self.model.stream_chat(**self.data)
else:
response = self.model.chat(**self.data)
results.raw = response
return results
async def _async_call(self, **kwargs) -> ModelResponse:

View file

@ -1,18 +1,17 @@
from typing import List, Dict
from typing import List
from llama_index.core.data_structs import Node
from llama_index.core.schema import NodeWithScore # type: ignore
from llama_index.core.schema import NodeWithScore
from llama_index.postprocessor.dashscope_rerank import DashScopeRerank
from memory_scope.enumeration.model_enum import ModelEnum
from memory_scope.models import MODEL_REGISTRY
from memory_scope.models.base_model import BaseModel
from memory_scope.models.response import ModelResponse, ModelResponseGen
from llama_index.postprocessor.dashscope_rerank import DashScopeRerank
from memory_scope.enumeration.model_enum import ModelEnum
from memory_scope.models.response import ModelResponse
class LlamaIndexRerankModel(BaseModel):
model_type: ModelEnum = ModelEnum.RANK_MODEL
m_type: ModelEnum = ModelEnum.RANK_MODEL
MODEL_REGISTRY.batch_register([
DashScopeRerank
@ -22,44 +21,35 @@ class LlamaIndexRerankModel(BaseModel):
assert "query" in kwargs or "documents" in kwargs
query: str = kwargs.pop("query", "")
documents: List[str] = kwargs.pop("documents", [])
assert query and documents, f"query or documents is empty! query={query}, documents={len(documents)}"
# using -1.0 as dummy scores
nodes = [NodeWithScore(node=Node(text=doc), score=-1.0) for doc in documents]
nodes = [NodeWithScore(node=Node(text=doc), score=-1.0) for doc in documents]
self._get_documents_mapping(documents)
self.data = {
"nodes": nodes,
"query_str": query,
}
def after_call(self, model_response: ModelResponse, stream: bool = False, **kwargs) -> ModelResponse:
nodes = model_response.raw
ranks = []
for node in nodes:
def after_call(self, model_response: ModelResponse, **kwargs) -> ModelResponse:
if not model_response.rank_scores:
model_response.rank_scores = {}
for node in model_response.raw:
text = node.node.text
idx = self.documents_map[text]
ranks.append({idx: node.score})
model_response.rank_scores = ranks
model_response.rank_scores[idx] = node.score
return model_response
def _call(self, **kwargs) -> ModelResponse:
results = ModelResponse(model_type=self.model_type)
try:
response = self.model.postprocess_nodes(**self.data)
results.raw = response
results.status = True
except Exception as e:
results.details = e
results.status = False
return results
return ModelResponse(m_type=self.m_type, raw=self.model.postprocess_nodes(**self.data))
async def _async_call(self, **kwargs) -> ModelResponse:
raise NotImplementedError
def _get_documents_mapping(self, documents):
self.documents_map = {}
for idx, doc in enumerate(documents):
self.documents_map[doc] = idx

View file

@ -1,3 +1,4 @@
import json
from typing import Generator, List, Dict, Any
from pydantic import BaseModel, Field
@ -12,10 +13,10 @@ class ModelResponse(BaseModel):
embedding_results: List[List[float]] | List[float] = Field([], description="embedding vector")
rank_scores: List[Dict[int, float]] = Field([], description="The rank scores of each documents. "
"key: index, value: rank score")
rank_scores: Dict[int, float] = Field({}, description="The rank scores of each documents. "
"key: index, value: rank score")
model_type: ModelEnum = Field(ModelEnum.GENERATION_MODEL, description="One of LLM, EMB, RANK.")
m_type: ModelEnum = Field(ModelEnum.GENERATION_MODEL, description="One of LLM, EMB, RANK.")
status: bool = Field(True, description="Indicates whether the model call was successful.")
@ -24,5 +25,24 @@ class ModelResponse(BaseModel):
raw: Any = Field("", description="Raw response from model call")
def __str__(self, max_size=100, **kwargs):
result = {}
try:
all_dict = self.model_dump()
except Exception:
all_dict = self.dict()
for key, value in all_dict.items():
if key == "raw" or not value:
continue
if isinstance(value, str):
result[key] = value
elif isinstance(value, list | dict):
result[key] = f"{str(value)[:max_size]}... size={len(value)}"
elif isinstance(value, ModelEnum):
result[key] = value.value
return json.dumps(result, **kwargs)
ModelResponseGen = Generator[ModelResponse, None, None]

View file

@ -18,8 +18,13 @@ class Registry(object):
raise KeyError(f'{module_name} is already registered in {self.name}')
self.module_dict[module_name] = module
def batch_register(self, modules: List[Any]):
module_name_dict = {m.__name__: m for m in modules}
def batch_register(self, modules: List[Any] | Dict[str, Any]):
if isinstance(modules, list):
module_name_dict = {m.__name__: m for m in modules}
elif isinstance(modules, dict):
module_name_dict = modules
else:
raise NotImplementedError
self.module_dict.update(module_name_dict)
def get(self, module_name: str):

View file

@ -1,6 +1,9 @@
import asyncio
import unittest
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
class TestLLIEmbedding(unittest.TestCase):
"""Tests for LlamaIndexEmbeddingModel"""
@ -11,19 +14,21 @@ class TestLLIEmbedding(unittest.TestCase):
"clazz": "models.base_embedding_model"
}
self.emb = LlamaIndexEmbeddingModel(**config)
def test_single_embedding(self):
text = "您吃了吗?"
embs = self.emb.call(text=text)
result = self.emb.call(text=text)
print(result)
def test_batch_embedding(self):
texts = ["您吃了吗?",
"吃了吗您?"]
embs = self.emb.call(text=texts)
texts = ["您吃了吗?",
"吃了吗您?"]
result = self.emb.call(text=texts)
print(result)
async def test_async_embedding(self):
texts = ["您吃了吗?",
"吃了吗您?"]
def test_async_embedding(self):
texts = ["您吃了吗?",
"吃了吗您?"]
# 调用异步函数并等待其结果
embs = await self.emb.async_call(texts)
result = asyncio.run(self.emb.async_call(text=texts))
print(result)

View file

@ -1,6 +1,9 @@
import json
import unittest
from memory_scope.models.llama_index_rerank_model import LlamaIndexRerankModel
class TestLLIReRank(unittest.TestCase):
"""Tests for LlamaIndexRerankModel"""
@ -14,10 +17,10 @@ class TestLLIReRank(unittest.TestCase):
def test_rerank(self):
query = "吃啥?"
documents = ["您吃了吗?",
"吃了吗您?"]
embs = self.reranker.call(
stream=False,
documents=documents,
query=query)
print(embs)
documents = ["您吃了吗?",
"吃了吗您?"]
result = self.reranker.call(
stream=False,
documents=documents,
query=query)
print(result)