codekingpro/portable-devtools
114k
1import logging2import os3from typing import Any, Dict, List, Mapping, Optional4 5import requests6from langchain_core.callbacks import CallbackManagerForLLMRun7from langchain_core.language_models.llms import LLM8from pydantic import Field9 10logger = logging.getLogger(__name__)11 12 13class Baseten(LLM):14 """Baseten model15 16 This module allows using LLMs hosted on Baseten.17 18 The LLM deployed on Baseten must have the following properties:19 20 * Must accept input as a dictionary with the key "prompt"21 * May accept other input in the dictionary passed through with kwargs22 * Must return a string with the model output23 24 To use this module, you must:25 26 * Export your Baseten API key as the environment variable `BASETEN_API_KEY`27 * Get the model ID for your model from your Baseten dashboard28 * Identify the model deployment ("production" for all model library models)29 30 These code samples use31 [Mistral 7B Instruct](https://app.baseten.co/explore/mistral_7b_instruct)32 from Baseten's model library.33 34 Examples:35 .. code-block:: python36 37 from langchain_community.llms import Baseten38 # Production deployment39 mistral = Baseten(model="MODEL_ID", deployment="production")40 mistral("What is the Mistral wind?")41 42 .. code-block:: python43 44 from langchain_community.llms import Baseten45 # Development deployment46 mistral = Baseten(model="MODEL_ID", deployment="development")47 mistral("What is the Mistral wind?")48 49 .. code-block:: python50 51 from langchain_community.llms import Baseten52 # Other published deployment53 mistral = Baseten(model="MODEL_ID", deployment="DEPLOYMENT_ID")54 mistral("What is the Mistral wind?")55 """56 57 model: str58 deployment: str59 input: Dict[str, Any] = Field(default_factory=dict)60 model_kwargs: Dict[str, Any] = Field(default_factory=dict)61 62 @property63 def _identifying_params(self) -> Mapping[str, Any]:64 """Get the identifying parameters."""65 return {66 **{"model_kwargs": self.model_kwargs},67 }68 69 @property70 def _llm_type(self) -> str:71 """Return type of model."""72 return "baseten"73 74 def _call(75 self,76 prompt: str,77 stop: Optional[List[str]] = None,78 run_manager: Optional[CallbackManagerForLLMRun] = None,79 **kwargs: Any,80 ) -> str:81 baseten_api_key = os.environ["BASETEN_API_KEY"]82 model_id = self.model83 if self.deployment == "production":84 model_url = f"https://model-{model_id}.api.baseten.co/production/predict"85 elif self.deployment == "development":86 model_url = f"https://model-{model_id}.api.baseten.co/development/predict"87 else: # try specific deployment ID88 model_url = f"https://model-{model_id}.api.baseten.co/deployment/{self.deployment}/predict"89 response = requests.post(90 model_url,91 headers={"Authorization": f"Api-Key {baseten_api_key}"},92 json={"prompt": prompt, **kwargs},93 )94 return response.json()95 