codekingpro/portable-devtools
114k
1from typing import Any, Dict, List, Mapping, Optional, Tuple2 3import requests4from langchain_core.embeddings import Embeddings5from langchain_core.utils import get_from_dict_or_env6from pydantic import BaseModel, ConfigDict, model_validator7 8 9class MosaicMLInstructorEmbeddings(BaseModel, Embeddings):10 """MosaicML embedding service.11 12 To use, you should have the13 environment variable ``MOSAICML_API_TOKEN`` set with your API token, or pass14 it as a named parameter to the constructor.15 16 Example:17 .. code-block:: python18 19 from langchain_community.llms import MosaicMLInstructorEmbeddings20 endpoint_url = (21 "https://models.hosted-on.mosaicml.hosting/instructor-large/v1/predict"22 )23 mosaic_llm = MosaicMLInstructorEmbeddings(24 endpoint_url=endpoint_url,25 mosaicml_api_token="my-api-key"26 )27 """28 29 endpoint_url: str = (30 "https://models.hosted-on.mosaicml.hosting/instructor-xl/v1/predict"31 )32 """Endpoint URL to use."""33 embed_instruction: str = "Represent the document for retrieval: "34 """Instruction used to embed documents."""35 query_instruction: str = (36 "Represent the question for retrieving supporting documents: "37 )38 """Instruction used to embed the query."""39 retry_sleep: float = 1.040 """How long to try sleeping for if a rate limit is encountered"""41 42 mosaicml_api_token: Optional[str] = None43 44 model_config = ConfigDict(45 extra="forbid",46 )47 48 @model_validator(mode="before")49 @classmethod50 def validate_environment(cls, values: Dict) -> Any:51 """Validate that api key and python package exists in environment."""52 mosaicml_api_token = get_from_dict_or_env(53 values, "mosaicml_api_token", "MOSAICML_API_TOKEN"54 )55 values["mosaicml_api_token"] = mosaicml_api_token56 return values57 58 @property59 def _identifying_params(self) -> Mapping[str, Any]:60 """Get the identifying parameters."""61 return {"endpoint_url": self.endpoint_url}62 63 def _embed(64 self, input: List[Tuple[str, str]], is_retry: bool = False65 ) -> List[List[float]]:66 payload = {"inputs": input}67 68 # HTTP headers for authorization69 headers = {70 "Authorization": f"{self.mosaicml_api_token}",71 "Content-Type": "application/json",72 }73 74 # send request75 try:76 response = requests.post(self.endpoint_url, headers=headers, json=payload)77 except requests.exceptions.RequestException as e:78 raise ValueError(f"Error raised by inference endpoint: {e}")79 80 try:81 if response.status_code == 429:82 if not is_retry:83 import time84 85 time.sleep(self.retry_sleep)86 87 return self._embed(input, is_retry=True)88 89 raise ValueError(90 f"Error raised by inference API: rate limit exceeded.\nResponse: "91 f"{response.text}"92 )93 94 parsed_response = response.json()95 96 # The inference API has changed a couple of times, so we add some handling97 # to be robust to multiple response formats.98 if isinstance(parsed_response, dict):99 output_keys = ["data", "output", "outputs"]100 for key in output_keys:101 if key in parsed_response:102 output_item = parsed_response[key]103 break104 else:105 raise ValueError(106 f"No key data or output in response: {parsed_response}"107 )108 109 if isinstance(output_item, list) and isinstance(output_item[0], list):110 embeddings = output_item111 else:112 embeddings = [output_item]113 else:114 raise ValueError(f"Unexpected response type: {parsed_response}")115 116 except requests.exceptions.JSONDecodeError as e:117 raise ValueError(118 f"Error raised by inference API: {e}.\nResponse: {response.text}"119 )120 121 return embeddings122 123 def embed_documents(self, texts: List[str]) -> List[List[float]]:124 """Embed documents using a MosaicML deployed instructor embedding model.125 126 Args:127 texts: The list of texts to embed.128 129 Returns:130 List of embeddings, one for each text.131 """132 instruction_pairs = [(self.embed_instruction, text) for text in texts]133 embeddings = self._embed(instruction_pairs)134 return embeddings135 136 def embed_query(self, text: str) -> List[float]:137 """Embed a query using a MosaicML deployed instructor embedding model.138 139 Args:140 text: The text to embed.141 142 Returns:143 Embeddings for the text.144 """145 instruction_pair = (self.query_instruction, text)146 embedding = self._embed([instruction_pair])[0]147 return embedding148 