codekingpro/portable-devtools
114k
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 