Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
rerank.py160 linesDownload Raw Back to rerank
1from json import dumps2from typing import Optional3 4import httpx5from requests import post6from yarl import URL7 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 22 23class OAICompatRerankModel(RerankModel):24    """25    rerank model API is compatible with Jina rerank model API. So copy the JinaRerankModel class code here.26    we need enhance for llama.cpp , which return raw score, not normalize score 0~1.  It seems Dify need it27    """28 29    def _invoke(30        self,31        model: str,32        credentials: dict,33        query: str,34        docs: list[str],35        score_threshold: Optional[float] = None,36        top_n: Optional[int] = None,37        user: Optional[str] = None,38    ) -> RerankResult:39        """40        Invoke rerank model41 42        :param model: model name43        :param credentials: model credentials44        :param query: search query45        :param docs: docs for reranking46        :param score_threshold: score threshold47        :param top_n: top n documents to return48        :param user: unique user id49        :return: rerank result50        """51        if len(docs) == 0:52            return RerankResult(model=model, docs=[])53 54        server_url = credentials["endpoint_url"]55        model_name = model56 57        if not server_url:58            raise CredentialsValidateFailedError("server_url is required")59        if not model_name:60            raise CredentialsValidateFailedError("model_name is required")61 62        url = server_url63        headers = {"Authorization": f"Bearer {credentials.get('api_key')}", "Content-Type": "application/json"}64 65        # TODO: Do we need truncate docs to avoid llama.cpp return error?66 67        data = {"model": model_name, "query": query, "documents": docs, "top_n": top_n}68 69        try:70            response = post(str(URL(url) / "rerank"), headers=headers, data=dumps(data), timeout=60)71            response.raise_for_status()72            results = response.json()73 74            rerank_documents = []75            scores = [result["relevance_score"] for result in results["results"]]76 77            # Min-Max Normalization: Normalize scores to 0 ~ 1.0 range78            min_score = min(scores)79            max_score = max(scores)80            score_range = max_score - min_score if max_score != min_score else 1.0  # Avoid division by zero81 82            for result in results["results"]:83                index = result["index"]84 85                # Retrieve document text (fallback if llama.cpp rerank doesn't return it)86                text = result.get("document", {}).get("text", docs[index])87 88                # Normalize the score89                normalized_score = (result["relevance_score"] - min_score) / score_range90 91                # Create RerankDocument object with normalized score92                rerank_document = RerankDocument(93                    index=index,94                    text=text,95                    score=normalized_score,96                )97 98                # Apply threshold (if defined)99                if score_threshold is None or normalized_score >= score_threshold:100                    rerank_documents.append(rerank_document)101 102            # Sort rerank_documents by normalized score in descending order103            rerank_documents.sort(key=lambda doc: doc.score, reverse=True)104 105            return RerankResult(model=model, docs=rerank_documents)106 107        except httpx.HTTPStatusError as e:108            raise InvokeServerUnavailableError(str(e))109 110    def validate_credentials(self, model: str, credentials: dict) -> None:111        """112        Validate model credentials113 114        :param model: model name115        :param credentials: model credentials116        :return:117        """118        try:119            self._invoke(120                model=model,121                credentials=credentials,122                query="What is the capital of the United States?",123                docs=[124                    "Carson City is the capital city of the American state of Nevada. At the 2010 United States "125                    "Census, Carson City had a population of 55,274.",126                    "The Commonwealth of the Northern Mariana Islands is a group of islands in the Pacific Ocean that "127                    "are a political division controlled by the United States. Its capital is Saipan.",128                ],129                score_threshold=0.8,130            )131        except Exception as ex:132            raise CredentialsValidateFailedError(str(ex))133 134    @property135    def _invoke_error_mapping(self) -> dict[type[InvokeError], list[type[Exception]]]:136        """137        Map model invoke error to unified error138        """139        return {140            InvokeConnectionError: [httpx.ConnectError],141            InvokeServerUnavailableError: [httpx.RemoteProtocolError],142            InvokeRateLimitError: [],143            InvokeAuthorizationError: [httpx.HTTPStatusError],144            InvokeBadRequestError: [httpx.RequestError],145        }146 147    def get_customizable_model_schema(self, model: str, credentials: dict) -> AIModelEntity:148        """149        generate custom model entities from credentials150        """151        entity = AIModelEntity(152            model=model,153            label=I18nObject(en_US=model),154            model_type=ModelType.RERANK,155            fetch_from=FetchFrom.CUSTOMIZABLE_MODEL,156            model_properties={},157        )158 159        return entity160