update module names, add stream support for llm

This commit is contained in:
xianzhe.xxz 2024-06-24 18:18:43 +08:00
parent 610e8f0b85
commit ffc1c5261d
7 changed files with 86 additions and 60 deletions

View file

@ -4,29 +4,15 @@ from llama_index.embeddings.dashscope import DashScopeEmbedding
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
class LlamaIndexEmbeddingModel(BaseModel):
model_type: ModelEnum = ModelEnum.EMBEDDING_MODEL
class BaseEmbeddingModel(BaseModel):
MODEL_REGISTRY.batch_register([
DashScopeEmbedding,
])
def before_call(self, **kwargs) -> None:
pass
def after_call(self, model_response: ModelResponse | ModelResponseGen,
**kwargs) -> ModelResponse | ModelResponseGen:
pass
def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen:
pass
async def _async_call(self, **kwargs) -> ModelResponse:
pass
class LLIEmbedding(BaseEmbeddingModel):
def before_call(self, **kwargs):
text: str | List[str] = kwargs.pop("text", "")
if isinstance(text, str):
@ -39,7 +25,7 @@ class LLIEmbedding(BaseEmbeddingModel):
return model_response
def _call(self, **kwargs) -> ModelResponse:
results = ModelResponse()
results = ModelResponse(model_type=self.model_type)
try:
response = self.model.get_text_embedding_batch(**self.data)
results.raw = response
@ -54,9 +40,9 @@ class LLIEmbedding(BaseEmbeddingModel):
:param kwargs:
:return:
"""
results = ModelResponse()
results = ModelResponse(model_type=self.model_type)
try:
response = self.model.aget_text_embedding_batch(**self.data)
response = await self.model.aget_text_embedding_batch(**self.data)
results.raw = response
results.status = True
except Exception as e:

View file

@ -44,20 +44,30 @@ class LlamaIndexGenerationModel(BaseModel):
def after_call(
self, model_response: ModelResponse, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen:
call_result = model_response.raw
if isinstance(call_result, CompletionResponse):
content = call_result.text
elif isinstance(call_result, ChatResponse):
content = call_result.message.content
if stream:
def gen() -> ModelResponseGen:
content = ""
for response in call_result:
delta = response.delta
content += delta
model_response.text = content
model_response.delta = delta
yield model_response
return gen()
else:
raise NotImplementedError
return ModelResponse(text=content,
model_type="LLM")
if isinstance(call_result, CompletionResponse):
content = call_result.text
elif isinstance(call_result, ChatResponse):
content = call_result.message.content
else:
raise NotImplementedError
model_response.text = content
return model_response
def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen:
assert "prompt" in self.data or "messages" in self.data
results = ModelResponse()
results = ModelResponse(model_type=self.model_type)
try:
if 'prompt' in self.data:
if stream:
@ -77,4 +87,4 @@ class LlamaIndexGenerationModel(BaseModel):
return results
async def _async_call(self, **kwargs) -> ModelResponse:
pass
raise NotImplementedError

View file

@ -7,28 +7,17 @@ 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
class BaseRankModel(BaseModel):
class LlamaIndexRerankModel(BaseModel):
model_type: ModelEnum = ModelEnum.RANK_MODEL
MODEL_REGISTRY.batch_register([
DashScopeRerank
])
def before_call(self, **kwargs) -> None:
pass
def after_call(self, **kwargs) -> ModelResponse:
pass
def _call(self, stream: bool = False, **kwargs) -> ModelResponse:
pass
async def _async_call(self, **kwargs) -> ModelResponse:
pass
class LLIReRank(BaseRankModel):
def before_call(self, **kwargs) -> None:
assert "query" in kwargs or "documents" in kwargs
query: str = kwargs.pop("query", "")
@ -37,7 +26,8 @@ class LLIReRank(BaseRankModel):
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=text), score=-1.0) for text in documents]
nodes = [NodeWithScore(node=Node(text=doc), score=-1.0) for doc in documents]
self._get_documents_mapping(documents)
self.data = {
"nodes": nodes,
@ -46,15 +36,16 @@ class LLIReRank(BaseRankModel):
def after_call(self, model_response: ModelResponse, stream: bool = False, **kwargs) -> ModelResponse:
nodes = model_response.raw
ranks = list()
ranks = []
for node in nodes:
ranks.append(dict(relevance_score=node.score,
document=node.node.text))
results = ModelResponse(rank_scores=ranks)
return results
text = node.node.text
idx = self.documents_map[text]
ranks.append({idx: node.score})
model_response.rank_scores = ranks
return model_response
def _call(self, **kwargs) -> ModelResponse:
results = ModelResponse()
results = ModelResponse(model_type=self.model_type)
try:
response = self.model.postprocess_nodes(**self.data)
results.raw = response
@ -63,3 +54,12 @@ class LLIReRank(BaseRankModel):
results.details = e
results.status = False
return results
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

@ -20,7 +20,8 @@ class ModelResponse(BaseModel):
details: str = Field("", description="The details information for model call, "
"usually for storage of raw response or failure messages.")
raw: Any = Field("", description="raw response from model call")
raw: Any = Field("", description="Raw response from model call")
delta: str = Field("", description="New text that just streamed in (only used when streaming)")
ModelResponseGen = Generator[ModelResponse, None, None]

View file

@ -1,5 +1,5 @@
import unittest
from memory_scope.models.base_embedding_model import LLIEmbedding
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
class TestLLIEmbedding(unittest.TestCase):
"""Tests for LLIEmbedding"""
@ -10,7 +10,7 @@ class TestLLIEmbedding(unittest.TestCase):
"model_name": "text-embedding-v2",
"clazz": "models.base_embedding_model"
}
self.emb = LLIEmbedding(**config)
self.emb = LlamaIndexEmbeddingModel(**config)
def test_single_embedding(self):
text = "您吃了吗?"

View file

@ -30,3 +30,31 @@ class TestLLILLM(unittest.TestCase):
messages=messages
)
print(ans.text)
def test_llm_prompt_stream(self):
prompt = "你如何看待黄金上涨?"
ans = self.llm.call(
stream=True,
prompt=prompt
)
import sys
import time
for a in ans:
sys.stdout.write(a.delta)
sys.stdout.flush()
time.sleep(0.1)
@unittest.skip('tmp')
def test_llm_messages(self):
messages = [{"role": "system", "content": "you are a helpful assistant."},
{"role": "user", "content": "你如何看待黄金上涨?"}]
ans = self.llm.call(
stream=True,
messages=messages
)
import sys
import time
for a in ans:
sys.stdout.write(a.delta)
sys.stdout.flush()
time.sleep(0.1)

View file

@ -1,5 +1,5 @@
import unittest
from memory_scope.models.base_rank_model import LLIReRank
from memory_scope.models.llama_index_rerank_model import LlamaIndexRerankModel
class TestLLIReRank(unittest.TestCase):
"""Tests for LLIEmbedding"""
@ -8,9 +8,9 @@ class TestLLIReRank(unittest.TestCase):
config = {
"method_type": "DashScopeRerank",
"model_name": "gte-rerank",
"clazz": "models.base_rank_model"
"clazz": "models.llama_index_rerank_model"
}
self.reranker = LLIReRank(**config)
self.reranker = LlamaIndexRerankModel(**config)
def test_rerank(self):
query = "吃啥?"
@ -20,3 +20,4 @@ class TestLLIReRank(unittest.TestCase):
stream=False,
documents=documents,
query=query)
print(embs)