Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
oci_generative_ai.py233 linesDownload Raw Back to embeddings
1from enum import Enum2from typing import TYPE_CHECKING, Any, Dict, Iterator, List, Mapping, Optional3 4from langchain_core.embeddings import Embeddings5from langchain_core.utils import pre_init6from pydantic import BaseModel, ConfigDict7 8if TYPE_CHECKING:9    import oci10 11CUSTOM_ENDPOINT_PREFIX = "ocid1.generativeaiendpoint"12 13 14class OCIAuthType(Enum):15    """OCI authentication types as enumerator."""16 17    API_KEY = 118    SECURITY_TOKEN = 219    INSTANCE_PRINCIPAL = 320    RESOURCE_PRINCIPAL = 421 22 23class OCIGenAIEmbeddings(BaseModel, Embeddings):24    """OCI embedding models.25 26    To authenticate, the OCI client uses the methods described in27    https://docs.oracle.com/en-us/iaas/Content/API/Concepts/sdk_authentication_methods.htm28 29    The authentifcation method is passed through auth_type and should be one of:30    API_KEY (default), SECURITY_TOKEN, INSTANCE_PRINCIPLE, RESOURCE_PRINCIPLE31 32    Make sure you have the required policies (profile/roles) to33    access the OCI Generative AI service. If a specific config profile is used,34    you must pass the name of the profile (~/.oci/config) through auth_profile.35    If a specific config file location is used, you must pass36    the file location where profile name configs present37    through auth_file_location38 39    To use, you must provide the compartment id40    along with the endpoint url, and model id41    as named parameters to the constructor.42 43    Example:44        .. code-block:: python45 46            from langchain_classic.embeddings import OCIGenAIEmbeddings47 48            embeddings = OCIGenAIEmbeddings(49                model_id="MY_EMBEDDING_MODEL",50                service_endpoint="https://inference.generativeai.us-chicago-1.oci.oraclecloud.com",51                compartment_id="MY_OCID"52            )53    """54 55    client: Any = None  #: :meta private:56 57    service_models: Any = None  #: :meta private:58 59    auth_type: Optional[str] = "API_KEY"60    """Authentication type, could be 61 62    API_KEY, 63    SECURITY_TOKEN, 64    INSTANCE_PRINCIPLE, 65    RESOURCE_PRINCIPLE66    67    If not specified, API_KEY will be used68    """69 70    auth_profile: Optional[str] = "DEFAULT"71    """The name of the profile in ~/.oci/config72    If not specified , DEFAULT will be used 73    """74 75    auth_file_location: Optional[str] = "~/.oci/config"76    """Path to the config file.77    If not specified, ~/.oci/config will be used78    """79 80    model_id: Optional[str] = None81    """Id of the model to call, e.g., cohere.embed-english-light-v2.0"""82 83    model_kwargs: Optional[Dict] = None84    """Keyword arguments to pass to the model"""85 86    service_endpoint: Optional[str] = None87    """service endpoint url"""88 89    compartment_id: Optional[str] = None90    """OCID of compartment"""91 92    truncate: Optional[str] = "END"93    """Truncate embeddings that are too long from start or end ("NONE"|"START"|"END")"""94 95    batch_size: int = 9696    """Batch size of OCI GenAI embedding requests. OCI GenAI may handle up to 96 texts97     per request"""98 99    model_config = ConfigDict(extra="forbid", protected_namespaces=())100 101    @pre_init102    def validate_environment(cls, values: Dict) -> Dict:  # pylint: disable=no-self-argument103        """Validate that OCI config and python package exists in environment."""104 105        # Skip creating new client if passed in constructor106        if values["client"] is not None:107            return values108 109        try:110            import oci111 112            client_kwargs = {113                "config": {},114                "signer": None,115                "service_endpoint": values["service_endpoint"],116                "retry_strategy": oci.retry.DEFAULT_RETRY_STRATEGY,117                "timeout": (10, 240),  # default timeout config for OCI Gen AI service118            }119 120            if values["auth_type"] == OCIAuthType(1).name:121                client_kwargs["config"] = oci.config.from_file(122                    file_location=values["auth_file_location"],123                    profile_name=values["auth_profile"],124                )125                client_kwargs.pop("signer", None)126            elif values["auth_type"] == OCIAuthType(2).name:127 128                def make_security_token_signer(129                    oci_config: dict[str, Any],130                ) -> "oci.auth.signers.SecurityTokenSigner":131                    pk = oci.signer.load_private_key_from_file(132                        oci_config.get("key_file"), None133                    )134                    with open(135                        str(oci_config.get("security_token_file")), encoding="utf-8"136                    ) as f:137                        st_string = f.read()138                    return oci.auth.signers.SecurityTokenSigner(st_string, pk)139 140                client_kwargs["config"] = oci.config.from_file(141                    file_location=values["auth_file_location"],142                    profile_name=values["auth_profile"],143                )144                client_kwargs["signer"] = make_security_token_signer(145                    oci_config=client_kwargs["config"]146                )147            elif values["auth_type"] == OCIAuthType(3).name:148                client_kwargs["signer"] = (149                    oci.auth.signers.InstancePrincipalsSecurityTokenSigner()150                )151            elif values["auth_type"] == OCIAuthType(4).name:152                client_kwargs["signer"] = (153                    oci.auth.signers.get_resource_principals_signer()154                )155            else:156                raise ValueError("Please provide valid value to auth_type")157 158            values["client"] = oci.generative_ai_inference.GenerativeAiInferenceClient(159                **client_kwargs160            )161 162        except ImportError as ex:163            raise ImportError(164                "Could not import oci python package. "165                "Please make sure you have the oci package installed."166            ) from ex167        except Exception as e:168            raise ValueError(169                """Could not authenticate with OCI client.170                If INSTANCE_PRINCIPAL or RESOURCE_PRINCIPAL is used,171                please check the specified172                auth_profile, auth_file_location and auth_type are valid.""",173                e,174            ) from e175 176        return values177 178    @property179    def _identifying_params(self) -> Mapping[str, Any]:180        """Get the identifying parameters."""181        _model_kwargs = self.model_kwargs or {}182        return {183            **{"model_kwargs": _model_kwargs},184        }185 186    def embed_documents(self, texts: List[str]) -> List[List[float]]:187        """Call out to OCIGenAI's embedding endpoint.188 189        Args:190            texts: The list of texts to embed.191 192        Returns:193            List of embeddings, one for each text.194        """195        from oci.generative_ai_inference import models196 197        if not self.model_id:198            raise ValueError("Model ID is required to embed documents")199 200        if self.model_id.startswith(CUSTOM_ENDPOINT_PREFIX):201            serving_mode = models.DedicatedServingMode(endpoint_id=self.model_id)202        else:203            serving_mode = models.OnDemandServingMode(model_id=self.model_id)204 205        embeddings = []206 207        def split_texts() -> Iterator[List[str]]:208            for i in range(0, len(texts), self.batch_size):209                yield texts[i : i + self.batch_size]210 211        for chunk in split_texts():212            invocation_obj = models.EmbedTextDetails(213                serving_mode=serving_mode,214                compartment_id=self.compartment_id,215                truncate=self.truncate,216                inputs=chunk,217            )218            response = self.client.embed_text(invocation_obj)219            embeddings.extend(response.data.embeddings)220 221        return embeddings222 223    def embed_query(self, text: str) -> List[float]:224        """Call out to OCIGenAI's embedding endpoint.225 226        Args:227            text: The text to embed.228 229        Returns:230            Embeddings for the text.231        """232        return self.embed_documents([text])[0]233 
codekingpro/portable-devtools · Team Ai