mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-06 02:48:22 +00:00
[dev] delete asyncio from base model
This commit is contained in:
parent
160652b600
commit
322f48117f
8 changed files with 132 additions and 112 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue