Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
baseten.py95 linesDownload Raw Back to llms
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 
codekingpro/portable-devtools · Team Ai