Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
rerank.py174 linesDownload Raw Back to rerank
1import json2import logging3import operator4from typing import Any, Optional5 6import boto37 8from core.model_runtime.entities.common_entities import I18nObject9from core.model_runtime.entities.model_entities import AIModelEntity, FetchFrom, ModelType10from core.model_runtime.entities.rerank_entities import RerankDocument, RerankResult11from core.model_runtime.errors.invoke import (12    InvokeAuthorizationError,13    InvokeBadRequestError,14    InvokeConnectionError,15    InvokeError,16    InvokeRateLimitError,17    InvokeServerUnavailableError,18)19from core.model_runtime.errors.validate import CredentialsValidateFailedError20from core.model_runtime.model_providers.__base.rerank_model import RerankModel21 22logger = logging.getLogger(__name__)23 24 25class SageMakerRerankModel(RerankModel):26    """27    Model class for SageMaker rerank model.28    """29 30    sagemaker_client: Any = None31 32    def _sagemaker_rerank(self, query_input: str, docs: list[str], rerank_endpoint: str):33        inputs = [query_input] * len(docs)34        response_model = self.sagemaker_client.invoke_endpoint(35            EndpointName=rerank_endpoint,36            Body=json.dumps({"inputs": inputs, "docs": docs}),37            ContentType="application/json",38        )39        json_str = response_model["Body"].read().decode("utf8")40        json_obj = json.loads(json_str)41        scores = json_obj["scores"]42        return scores if isinstance(scores, list) else [scores]43 44    def _invoke(45        self,46        model: str,47        credentials: dict,48        query: str,49        docs: list[str],50        score_threshold: Optional[float] = None,51        top_n: Optional[int] = None,52        user: Optional[str] = None,53    ) -> RerankResult:54        """55        Invoke rerank model56 57        :param model: model name58        :param credentials: model credentials59        :param query: search query60        :param docs: docs for reranking61        :param score_threshold: score threshold62        :param top_n: top n63        :param user: unique user id64        :return: rerank result65        """66        line = 067        try:68            if len(docs) == 0:69                return RerankResult(model=model, docs=docs)70 71            line = 172            if not self.sagemaker_client:73                access_key = credentials.get("aws_access_key_id")74                secret_key = credentials.get("aws_secret_access_key")75                aws_region = credentials.get("aws_region")76                if aws_region:77                    if access_key and secret_key:78                        self.sagemaker_client = boto3.client(79                            "sagemaker-runtime",80                            aws_access_key_id=access_key,81                            aws_secret_access_key=secret_key,82                            region_name=aws_region,83                        )84                    else:85                        self.sagemaker_client = boto3.client("sagemaker-runtime", region_name=aws_region)86                else:87                    self.sagemaker_client = boto3.client("sagemaker-runtime")88 89            line = 290 91            sagemaker_endpoint = credentials.get("sagemaker_endpoint")92            candidate_docs = []93 94            scores = self._sagemaker_rerank(query, docs, sagemaker_endpoint)95            for idx in range(len(scores)):96                candidate_docs.append({"content": docs[idx], "score": scores[idx]})97 98            sorted(candidate_docs, key=operator.itemgetter("score"), reverse=True)99 100            line = 3101            rerank_documents = []102            for idx, result in enumerate(candidate_docs):103                rerank_document = RerankDocument(104                    index=idx, text=result.get("content"), score=result.get("score", -100.0)105                )106 107                if score_threshold is not None:108                    if rerank_document.score >= score_threshold:109                        rerank_documents.append(rerank_document)110                else:111                    rerank_documents.append(rerank_document)112 113            return RerankResult(model=model, docs=rerank_documents)114 115        except Exception as e:116            logger.exception(f"Exception {e}, line : {line}")117 118    def validate_credentials(self, model: str, credentials: dict) -> None:119        """120        Validate model credentials121 122        :param model: model name123        :param credentials: model credentials124        :return:125        """126        try:127            self._invoke(128                model=model,129                credentials=credentials,130                query="What is the capital of the United States?",131                docs=[132                    "Carson City is the capital city of the American state of Nevada. At the 2010 United States "133                    "Census, Carson City had a population of 55,274.",134                    "The Commonwealth of the Northern Mariana Islands is a group of islands in the Pacific Ocean that "135                    "are a political division controlled by the United States. Its capital is Saipan.",136                ],137                score_threshold=0.8,138            )139        except Exception as ex:140            raise CredentialsValidateFailedError(str(ex))141 142    @property143    def _invoke_error_mapping(self) -> dict[type[InvokeError], list[type[Exception]]]:144        """145        Map model invoke error to unified error146        The key is the error type thrown to the caller147        The value is the error type thrown by the model,148        which needs to be converted into a unified error type for the caller.149 150        :return: Invoke error mapping151        """152        return {153            InvokeConnectionError: [InvokeConnectionError],154            InvokeServerUnavailableError: [InvokeServerUnavailableError],155            InvokeRateLimitError: [InvokeRateLimitError],156            InvokeAuthorizationError: [InvokeAuthorizationError],157            InvokeBadRequestError: [InvokeBadRequestError, KeyError, ValueError],158        }159 160    def get_customizable_model_schema(self, model: str, credentials: dict) -> Optional[AIModelEntity]:161        """162        used to define customizable model schema163        """164        entity = AIModelEntity(165            model=model,166            label=I18nObject(en_US=model),167            fetch_from=FetchFrom.CUSTOMIZABLE_MODEL,168            model_type=ModelType.RERANK,169            model_properties={},170            parameter_rules=[],171        )172 173        return entity174