mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-07 03:00:27 +00:00
update module names, add stream support for llm
This commit is contained in:
parent
610e8f0b85
commit
ffc1c5261d
7 changed files with 86 additions and 60 deletions
|
|
@ -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:
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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 = "您吃了吗?"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
Loading…
Add table
Reference in a new issue