Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
nemo.py191 linesDownload Raw Back to embeddings
1from __future__ import annotations2 3import asyncio4import json5from typing import Any, Dict, List, Optional6 7import aiohttp8import requests9from langchain_core._api.deprecation import deprecated10from langchain_core.embeddings import Embeddings11from langchain_core.utils import pre_init12from pydantic import BaseModel13 14 15def is_endpoint_live(url: str, headers: Optional[dict], payload: Any) -> bool:16    """17    Check if an endpoint is live by sending a GET request to the specified URL.18 19    Args:20        url (str): The URL of the endpoint to check.21 22    Returns:23        bool: True if the endpoint is live (status code 200), False otherwise.24 25    Raises:26        Exception: If the endpoint returns a non-successful status code or if there is27            an error querying the endpoint.28    """29    try:30        response = requests.request("POST", url, headers=headers, data=payload)31 32        # Check if the status code is 200 (OK)33        if response.status_code == 200:34            return True35        else:36            # Raise an exception if the status code is not 20037            raise Exception(38                f"Endpoint returned a non-successful status code: "39                f"{response.status_code}"40            )41    except requests.exceptions.RequestException as e:42        # Handle any exceptions (e.g., connection errors)43        raise Exception(f"Error querying the endpoint: {e}")44 45 46@deprecated(47    since="0.0.37",48    removal="1.0.0",49    message=(50        "Directly instantiating a NeMoEmbeddings from langchain-community is "51        "deprecated. Please use langchain-nvidia-ai-endpoints NVIDIAEmbeddings "52        "interface."53    ),54)55class NeMoEmbeddings(BaseModel, Embeddings):56    """NeMo embedding models."""57 58    batch_size: int = 1659    model: str = "NV-Embed-QA-003"60    api_endpoint_url: str = "http://localhost:8088/v1/embeddings"61 62    @pre_init63    def validate_environment(cls, values: Dict) -> Dict:64        """Validate that the end point is alive using the values that are provided."""65 66        url = values["api_endpoint_url"]67        model = values["model"]68 69        # Optional: A minimal test payload and headers required by the endpoint70        headers = {"Content-Type": "application/json"}71        payload = json.dumps(72            {73                "input": "Hello World",74                "model": model,75                "input_type": "query",76            }77        )78 79        is_endpoint_live(url, headers, payload)80 81        return values82 83    async def _aembedding_func(84        self, session: Any, text: str, input_type: str85    ) -> List[float]:86        """Async call out to embedding endpoint.87 88        Args:89            text: The text to embed.90 91        Returns:92            Embeddings for the text.93        """94 95        headers = {"Content-Type": "application/json"}96 97        async with session.post(98            self.api_endpoint_url,99            json={"input": text, "model": self.model, "input_type": input_type},100            headers=headers,101        ) as response:102            response.raise_for_status()103            answer = await response.text()104            answer = json.loads(answer)105            return answer["data"][0]["embedding"]106 107    def _embedding_func(self, text: str, input_type: str) -> List[float]:108        """Call out to Cohere's embedding endpoint.109 110        Args:111            text: The text to embed.112 113        Returns:114            Embeddings for the text.115        """116 117        payload = json.dumps(118            {119                "input": text,120                "model": self.model,121                "input_type": input_type,122            }123        )124        headers = {"Content-Type": "application/json"}125 126        response = requests.request(127            "POST", self.api_endpoint_url, headers=headers, data=payload128        )129        response_json = json.loads(response.text)130        embedding = response_json["data"][0]["embedding"]131 132        return embedding133 134    def embed_documents(self, documents: List[str]) -> List[List[float]]:135        """Embed a list of document texts.136 137        Args:138            texts: The list of texts to embed.139 140        Returns:141            List of embeddings, one for each text.142        """143        return [self._embedding_func(text, input_type="passage") for text in documents]144 145    def embed_query(self, text: str) -> List[float]:146        return self._embedding_func(text, input_type="query")147 148    async def aembed_query(self, text: str) -> List[float]:149        """Call out to NeMo's embedding endpoint async for embedding query text.150 151        Args:152            text: The text to embed.153 154        Returns:155            Embedding for the text.156        """157 158        async with aiohttp.ClientSession() as session:159            embedding = await self._aembedding_func(session, text, "passage")160            return embedding161 162    async def aembed_documents(self, texts: List[str]) -> List[List[float]]:163        """Call out to NeMo's embedding endpoint async for embedding search docs.164 165        Args:166            texts: The list of texts to embed.167 168        Returns:169            List of embeddings, one for each text.170        """171        embeddings = []172 173        async with aiohttp.ClientSession() as session:174            for batch in range(0, len(texts), self.batch_size):175                text_batch = texts[batch : batch + self.batch_size]176 177                for text in text_batch:178                    # Create tasks for all texts in the batch179                    tasks = [180                        self._aembedding_func(session, text, "passage")181                        for text in text_batch182                    ]183 184                    # Run all tasks concurrently185                    batch_results = await asyncio.gather(*tasks)186 187                    # Extend the embeddings list with results from this batch188                    embeddings.extend(batch_results)189 190        return embeddings191 
codekingpro/portable-devtools · Team Ai