codekingpro/portable-devtools
114k
1from typing import Any, Dict, List, Optional2 3from langchain_core._api.deprecation import deprecated4from langchain_core.embeddings import Embeddings5from langchain_core.utils import get_from_dict_or_env6from pydantic import BaseModel, ConfigDict, model_validator7 8from langchain_community.llms.cohere import _create_retry_decorator9 10 11@deprecated(12 since="0.0.30",13 removal="1.0",14 alternative_import="langchain_cohere.CohereEmbeddings",15)16class CohereEmbeddings(BaseModel, Embeddings):17 """Cohere embedding models.18 19 To use, you should have the ``cohere`` python package installed, and the20 environment variable ``COHERE_API_KEY`` set with your API key or pass it21 as a named parameter to the constructor.22 23 Example:24 .. code-block:: python25 26 from langchain_community.embeddings import CohereEmbeddings27 cohere = CohereEmbeddings(28 model="embed-english-light-v3.0",29 cohere_api_key="my-api-key"30 )31 """32 33 client: Any = None #: :meta private:34 """Cohere client."""35 async_client: Any = None #: :meta private:36 """Cohere async client."""37 model: str = "embed-english-v2.0"38 """Model name to use."""39 40 truncate: Optional[str] = None41 """Truncate embeddings that are too long from start or end ("NONE"|"START"|"END")"""42 43 cohere_api_key: Optional[str] = None44 45 max_retries: int = 346 """Maximum number of retries to make when generating."""47 request_timeout: Optional[float] = None48 """Timeout in seconds for the Cohere API request."""49 user_agent: str = "langchain"50 """Identifier for the application making the request."""51 52 model_config = ConfigDict(53 extra="forbid",54 )55 56 @model_validator(mode="before")57 @classmethod58 def validate_environment(cls, values: Dict) -> Any:59 """Validate that api key and python package exists in environment."""60 cohere_api_key = get_from_dict_or_env(61 values, "cohere_api_key", "COHERE_API_KEY"62 )63 request_timeout = values.get("request_timeout")64 65 try:66 import cohere67 68 client_name = values["user_agent"]69 values["client"] = cohere.Client(70 cohere_api_key,71 timeout=request_timeout,72 client_name=client_name,73 )74 values["async_client"] = cohere.AsyncClient(75 cohere_api_key,76 timeout=request_timeout,77 client_name=client_name,78 )79 except ImportError:80 raise ImportError(81 "Could not import cohere python package. "82 "Please install it with `pip install cohere`."83 )84 return values85 86 def embed_with_retry(self, **kwargs: Any) -> Any:87 """Use tenacity to retry the embed call."""88 retry_decorator = _create_retry_decorator(self.max_retries)89 90 @retry_decorator91 def _embed_with_retry(**kwargs: Any) -> Any:92 return self.client.embed(**kwargs)93 94 return _embed_with_retry(**kwargs)95 96 def aembed_with_retry(self, **kwargs: Any) -> Any:97 """Use tenacity to retry the embed call."""98 retry_decorator = _create_retry_decorator(self.max_retries)99 100 @retry_decorator101 async def _embed_with_retry(**kwargs: Any) -> Any:102 return await self.async_client.embed(**kwargs)103 104 return _embed_with_retry(**kwargs)105 106 def embed(107 self, texts: List[str], *, input_type: Optional[str] = None108 ) -> List[List[float]]:109 embeddings = self.embed_with_retry(110 model=self.model,111 texts=texts,112 input_type=input_type,113 truncate=self.truncate,114 ).embeddings115 return [list(map(float, e)) for e in embeddings]116 117 async def aembed(118 self, texts: List[str], *, input_type: Optional[str] = None119 ) -> List[List[float]]:120 embeddings = (121 await self.aembed_with_retry(122 model=self.model,123 texts=texts,124 input_type=input_type,125 truncate=self.truncate,126 )127 ).embeddings128 return [list(map(float, e)) for e in embeddings]129 130 def embed_documents(self, texts: List[str]) -> List[List[float]]:131 """Embed a list of document texts.132 133 Args:134 texts: The list of texts to embed.135 136 Returns:137 List of embeddings, one for each text.138 """139 return self.embed(texts, input_type="search_document")140 141 async def aembed_documents(self, texts: List[str]) -> List[List[float]]:142 """Async call out to Cohere's embedding endpoint.143 144 Args:145 texts: The list of texts to embed.146 147 Returns:148 List of embeddings, one for each text.149 """150 return await self.aembed(texts, input_type="search_document")151 152 def embed_query(self, text: str) -> List[float]:153 """Call out to Cohere's embedding endpoint.154 155 Args:156 text: The text to embed.157 158 Returns:159 Embeddings for the text.160 """161 return self.embed([text], input_type="search_query")[0]162 163 async def aembed_query(self, text: str) -> List[float]:164 """Async call out to Cohere's embedding endpoint.165 166 Args:167 text: The text to embed.168 169 Returns:170 Embeddings for the text.171 """172 return (await self.aembed([text], input_type="search_query"))[0]173 