codekingpro/portable-devtools
114k
1from __future__ import annotations2 3import logging4from functools import cached_property5from typing import Any, Dict, List, Optional6 7from langchain_core._api.deprecation import deprecated8from langchain_core.embeddings import Embeddings9from langchain_core.utils import pre_init10from langchain_core.utils.pydantic import get_fields11from pydantic import BaseModel12 13logger = logging.getLogger(__name__)14 15MAX_BATCH_SIZE_CHARS = 100000016MAX_BATCH_SIZE_PARTS = 9017 18 19@deprecated(20 since="0.3.5",21 removal="1.0",22 alternative_import="langchain_gigachat.GigaChatEmbeddings",23)24class GigaChatEmbeddings(BaseModel, Embeddings):25 """GigaChat Embeddings models.26 27 Example:28 .. code-block:: python29 from langchain_community.embeddings.gigachat import GigaChatEmbeddings30 31 embeddings = GigaChatEmbeddings(32 credentials=..., scope=..., verify_ssl_certs=...33 )34 """35 36 base_url: Optional[str] = None37 """ Base API URL """38 auth_url: Optional[str] = None39 """ Auth URL """40 credentials: Optional[str] = None41 """ Auth Token """42 scope: Optional[str] = None43 """ Permission scope for access token """44 45 access_token: Optional[str] = None46 """ Access token for GigaChat """47 48 model: Optional[str] = None49 """Model name to use."""50 user: Optional[str] = None51 """ Username for authenticate """52 password: Optional[str] = None53 """ Password for authenticate """54 55 timeout: Optional[float] = 60056 """ Timeout for request. By default it works for long requests. """57 verify_ssl_certs: Optional[bool] = None58 """ Check certificates for all requests """59 60 ca_bundle_file: Optional[str] = None61 cert_file: Optional[str] = None62 key_file: Optional[str] = None63 key_file_password: Optional[str] = None64 # Support for connection to GigaChat through SSL certificates65 66 @cached_property67 def _client(self) -> Any:68 """Returns GigaChat API client"""69 import gigachat70 71 return gigachat.GigaChat(72 base_url=self.base_url,73 auth_url=self.auth_url,74 credentials=self.credentials,75 scope=self.scope,76 access_token=self.access_token,77 model=self.model,78 user=self.user,79 password=self.password,80 timeout=self.timeout,81 verify_ssl_certs=self.verify_ssl_certs,82 ca_bundle_file=self.ca_bundle_file,83 cert_file=self.cert_file,84 key_file=self.key_file,85 key_file_password=self.key_file_password,86 )87 88 @pre_init89 def validate_environment(cls, values: Dict) -> Dict:90 """Validate authenticate data in environment and python package is installed."""91 try:92 import gigachat # noqa: F40193 except ImportError:94 raise ImportError(95 "Could not import gigachat python package. "96 "Please install it with `pip install gigachat`."97 )98 fields = set(get_fields(cls).keys())99 diff = set(values.keys()) - fields100 if diff:101 logger.warning(f"Extra fields {diff} in GigaChat class")102 return values103 104 def embed_documents(self, texts: List[str]) -> List[List[float]]:105 """Embed documents using a GigaChat embeddings models.106 107 Args:108 texts: The list of texts to embed.109 110 Returns:111 List of embeddings, one for each text.112 """113 result: List[List[float]] = []114 size = 0115 local_texts = []116 embed_kwargs = {}117 if self.model is not None:118 embed_kwargs["model"] = self.model119 for text in texts:120 local_texts.append(text)121 size += len(text)122 if size > MAX_BATCH_SIZE_CHARS or len(local_texts) > MAX_BATCH_SIZE_PARTS:123 for embedding in self._client.embeddings(124 texts=local_texts, **embed_kwargs125 ).data:126 result.append(embedding.embedding)127 size = 0128 local_texts = []129 # Call for last iteration130 if local_texts:131 for embedding in self._client.embeddings(132 texts=local_texts, **embed_kwargs133 ).data:134 result.append(embedding.embedding)135 136 return result137 138 async def aembed_documents(self, texts: List[str]) -> List[List[float]]:139 """Embed documents using a GigaChat embeddings models.140 141 Args:142 texts: The list of texts to embed.143 144 Returns:145 List of embeddings, one for each text.146 """147 result: List[List[float]] = []148 size = 0149 local_texts = []150 embed_kwargs = {}151 if self.model is not None:152 embed_kwargs["model"] = self.model153 for text in texts:154 local_texts.append(text)155 size += len(text)156 if size > MAX_BATCH_SIZE_CHARS or len(local_texts) > MAX_BATCH_SIZE_PARTS:157 embeddings = await self._client.aembeddings(158 texts=local_texts, **embed_kwargs159 )160 for embedding in embeddings.data:161 result.append(embedding.embedding)162 size = 0163 local_texts = []164 # Call for last iteration165 if local_texts:166 embeddings = await self._client.aembeddings(167 texts=local_texts, **embed_kwargs168 )169 for embedding in embeddings.data:170 result.append(embedding.embedding)171 172 return result173 174 def embed_query(self, text: str) -> List[float]:175 """Embed a query using a GigaChat embeddings models.176 177 Args:178 text: The text to embed.179 180 Returns:181 Embeddings for the text.182 """183 return self.embed_documents(texts=[text])[0]184 185 async def aembed_query(self, text: str) -> List[float]:186 """Embed a query using a GigaChat embeddings models.187 188 Args:189 text: The text to embed.190 191 Returns:192 Embeddings for the text.193 """194 docs = await self.aembed_documents(texts=[text])195 return docs[0]196 