diff --git a/memory_scope/models/base_embedding_model.py b/memory_scope/models/llama_index_embedding_model.py similarity index 69% rename from memory_scope/models/base_embedding_model.py rename to memory_scope/models/llama_index_embedding_model.py index 73069983..d79520c7 100644 --- a/memory_scope/models/base_embedding_model.py +++ b/memory_scope/models/llama_index_embedding_model.py @@ -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: diff --git a/memory_scope/models/llama_index_generation_model.py b/memory_scope/models/llama_index_generation_model.py index 0cbfcd47..e4a61f24 100644 --- a/memory_scope/models/llama_index_generation_model.py +++ b/memory_scope/models/llama_index_generation_model.py @@ -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 diff --git a/memory_scope/models/base_rank_model.py b/memory_scope/models/llama_index_rerank_model.py similarity index 65% rename from memory_scope/models/base_rank_model.py rename to memory_scope/models/llama_index_rerank_model.py index cfe5420d..511dbed5 100644 --- a/memory_scope/models/base_rank_model.py +++ b/memory_scope/models/llama_index_rerank_model.py @@ -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 + diff --git a/memory_scope/models/response.py b/memory_scope/models/response.py index 9a506293..72e98a44 100644 --- a/memory_scope/models/response.py +++ b/memory_scope/models/response.py @@ -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] diff --git a/tests/models/test_models_lli_embedding.py b/tests/models/test_models_lli_embedding.py index 6c380ba6..ecc2d401 100644 --- a/tests/models/test_models_lli_embedding.py +++ b/tests/models/test_models_lli_embedding.py @@ -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 = "您吃了吗?" diff --git a/tests/models/test_models_lli_llm.py b/tests/models/test_models_lli_generation.py similarity index 52% rename from tests/models/test_models_lli_llm.py rename to tests/models/test_models_lli_generation.py index 78c7e143..bc9fdfc5 100644 --- a/tests/models/test_models_lli_llm.py +++ b/tests/models/test_models_lli_generation.py @@ -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) \ No newline at end of file diff --git a/tests/models/test_models_lli_rerank.py b/tests/models/test_models_lli_rerank.py index 55dbe9bd..fbfd30ae 100644 --- a/tests/models/test_models_lli_rerank.py +++ b/tests/models/test_models_lli_rerank.py @@ -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) \ No newline at end of file