Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
text_embedding.py171 linesDownload Raw Back to text_embedding
1import time2from json import JSONDecodeError, dumps3from typing import Optional4 5import requests6 7from core.entities.embedding_type import EmbeddingInputType8from core.model_runtime.entities.common_entities import I18nObject9from core.model_runtime.entities.model_entities import AIModelEntity, FetchFrom, ModelPropertyKey, ModelType, PriceType10from core.model_runtime.entities.text_embedding_entities import EmbeddingUsage, TextEmbeddingResult11from 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.text_embedding_model import TextEmbeddingModel21 22 23class MixedBreadTextEmbeddingModel(TextEmbeddingModel):24    """25    Model class for MixedBread text embedding model.26    """27 28    api_base: str = "https://api.mixedbread.ai/v1"29 30    def _invoke(31        self,32        model: str,33        credentials: dict,34        texts: list[str],35        user: Optional[str] = None,36        input_type: EmbeddingInputType = EmbeddingInputType.DOCUMENT,37    ) -> TextEmbeddingResult:38        """39        Invoke text embedding model40 41        :param model: model name42        :param credentials: model credentials43        :param texts: texts to embed44        :param user: unique user id45        :param input_type: input type46        :return: embeddings result47        """48        api_key = credentials["api_key"]49        if not api_key:50            raise CredentialsValidateFailedError("api_key is required")51 52        base_url = credentials.get("base_url", self.api_base)53        base_url = base_url.removesuffix("/")54 55        url = base_url + "/embeddings"56        headers = {"Authorization": "Bearer " + api_key, "Content-Type": "application/json"}57 58        data = {"model": model, "input": texts}59 60        try:61            response = requests.post(url, headers=headers, data=dumps(data))62        except Exception as e:63            raise InvokeConnectionError(str(e))64 65        if response.status_code != 200:66            try:67                resp = response.json()68                msg = resp["detail"]69                if response.status_code == 401:70                    raise InvokeAuthorizationError(msg)71                elif response.status_code == 429:72                    raise InvokeRateLimitError(msg)73                elif response.status_code == 500:74                    raise InvokeServerUnavailableError(msg)75                else:76                    raise InvokeBadRequestError(msg)77            except JSONDecodeError as e:78                raise InvokeServerUnavailableError(79                    f"Failed to convert response to json: {e} with text: {response.text}"80                )81 82        try:83            resp = response.json()84            embeddings = resp["data"]85            usage = resp["usage"]86        except Exception as e:87            raise InvokeServerUnavailableError(f"Failed to convert response to json: {e} with text: {response.text}")88 89        usage = self._calc_response_usage(model=model, credentials=credentials, tokens=usage["total_tokens"])90 91        result = TextEmbeddingResult(92            model=model, embeddings=[[float(data) for data in x["embedding"]] for x in embeddings], usage=usage93        )94 95        return result96 97    def get_num_tokens(self, model: str, credentials: dict, texts: list[str]) -> int:98        """99        Get number of tokens for given prompt messages100 101        :param model: model name102        :param credentials: model credentials103        :param texts: texts to embed104        :return:105        """106        return sum(self._get_num_tokens_by_gpt2(text) for text in texts)107 108    def validate_credentials(self, model: str, credentials: dict) -> None:109        """110        Validate model credentials111 112        :param model: model name113        :param credentials: model credentials114        :return:115        """116        try:117            self._invoke(model=model, credentials=credentials, texts=["ping"])118        except Exception as e:119            raise CredentialsValidateFailedError(f"Credentials validation failed: {e}")120 121    @property122    def _invoke_error_mapping(self) -> dict[type[InvokeError], list[type[Exception]]]:123        return {124            InvokeConnectionError: [InvokeConnectionError],125            InvokeServerUnavailableError: [InvokeServerUnavailableError],126            InvokeRateLimitError: [InvokeRateLimitError],127            InvokeAuthorizationError: [InvokeAuthorizationError],128            InvokeBadRequestError: [KeyError, InvokeBadRequestError],129        }130 131    def _calc_response_usage(self, model: str, credentials: dict, tokens: int) -> EmbeddingUsage:132        """133        Calculate response usage134 135        :param model: model name136        :param credentials: model credentials137        :param tokens: input tokens138        :return: usage139        """140        # get input price info141        input_price_info = self.get_price(142            model=model, credentials=credentials, price_type=PriceType.INPUT, tokens=tokens143        )144 145        # transform usage146        usage = EmbeddingUsage(147            tokens=tokens,148            total_tokens=tokens,149            unit_price=input_price_info.unit_price,150            price_unit=input_price_info.unit,151            total_price=input_price_info.total_amount,152            currency=input_price_info.currency,153            latency=time.perf_counter() - self.started_at,154        )155 156        return usage157 158    def get_customizable_model_schema(self, model: str, credentials: dict) -> AIModelEntity:159        """160        generate custom model entities from credentials161        """162        entity = AIModelEntity(163            model=model,164            label=I18nObject(en_US=model),165            model_type=ModelType.TEXT_EMBEDDING,166            fetch_from=FetchFrom.CUSTOMIZABLE_MODEL,167            model_properties={ModelPropertyKey.CONTEXT_SIZE: int(credentials.get("context_size", "512"))},168        )169 170        return entity171