fix lli embedding

This commit is contained in:
xianzhe.xxz 2024-06-24 19:17:56 +08:00
parent a1c67d9e40
commit 160652b600
4 changed files with 5 additions and 4 deletions

View file

@ -22,6 +22,8 @@ class LlamaIndexEmbeddingModel(BaseModel):
def after_call(self, model_response: ModelResponse, **kwargs) -> ModelResponse:
embeddings = model_response.raw
model_response.embedding_results = embeddings
if len(embeddings) == 1:
model_response.embedding_results = embeddings[0]
return model_response
def _call(self, **kwargs) -> ModelResponse:

View file

@ -2,7 +2,7 @@ import unittest
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
class TestLLIEmbedding(unittest.TestCase):
"""Tests for LLIEmbedding"""
"""Tests for LlamaIndexEmbeddingModel"""
def setUp(self):
config = {

View file

@ -4,7 +4,7 @@ from memory_scope.models.llama_index_generation_model import LlamaIndexGeneratio
class TestLLILLM(unittest.TestCase):
"""Tests for LLIEmbedding"""
"""Tests for LlamaIndexGenerationModel"""
def setUp(self):
config = {
@ -44,7 +44,6 @@ class TestLLILLM(unittest.TestCase):
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": "你如何看待黄金上涨?"}]

View file

@ -2,7 +2,7 @@ import unittest
from memory_scope.models.llama_index_rerank_model import LlamaIndexRerankModel
class TestLLIReRank(unittest.TestCase):
"""Tests for LLIEmbedding"""
"""Tests for LlamaIndexRerankModel"""
def setUp(self):
config = {