codekingpro/portable-devtools
114k
1from typing import Any, List, Optional2 3import requests4from langchain_core.embeddings import Embeddings5from langchain_core.utils import (6 secret_from_env,7)8from pydantic import (9 BaseModel,10 ConfigDict,11 Field,12 SecretStr,13 model_validator,14)15from requests import RequestException16from typing_extensions import Self17 18BAICHUAN_API_URL: str = "https://api.baichuan-ai.com/v1/embeddings"19 20# BaichuanTextEmbeddings is an embedding model provided by Baichuan Inc. (https://www.baichuan-ai.com/home).21# As of today (Jan 25th, 2024) BaichuanTextEmbeddings ranks #1 in C-MTEB22# (Chinese Multi-Task Embedding Benchmark) leaderboard.23# Leaderboard (Under Overall -> Chinese section): https://huggingface.co/spaces/mteb/leaderboard24 25# Official Website: https://platform.baichuan-ai.com/docs/text-Embedding26# An API-key is required to use this embedding model. You can get one by registering27# at https://platform.baichuan-ai.com/docs/text-Embedding.28# BaichuanTextEmbeddings support 512 token window and produces vectors with29# 1024 dimensions.30 31 32# NOTE!! BaichuanTextEmbeddings only supports Chinese text embedding.33# Multi-language support is coming soon.34class BaichuanTextEmbeddings(BaseModel, Embeddings):35 """Baichuan Text Embedding models.36 37 Setup:38 To use, you should set the environment variable ``BAICHUAN_API_KEY`` to39 your API key or pass it as a named parameter to the constructor.40 41 .. code-block:: bash42 43 export BAICHUAN_API_KEY="your-api-key"44 45 Instantiate:46 .. code-block:: python47 48 from langchain_community.embeddings import BaichuanTextEmbeddings49 50 embeddings = BaichuanTextEmbeddings()51 52 Embed:53 .. code-block:: python54 55 # embed the documents56 vectors = embeddings.embed_documents([text1, text2, ...])57 58 # embed the query59 vectors = embeddings.embed_query(text)60 """ # noqa: E50161 62 session: Any = None #: :meta private:63 model_name: str = Field(default="Baichuan-Text-Embedding", alias="model")64 """The model used to embed the documents."""65 baichuan_api_key: SecretStr = Field(66 alias="api_key",67 default_factory=secret_from_env(["BAICHUAN_API_KEY", "BAICHUAN_AUTH_TOKEN"]),68 )69 """Automatically inferred from env var `BAICHUAN_API_KEY` if not provided."""70 chunk_size: int = 1671 """Chunk size when multiple texts are input"""72 73 model_config = ConfigDict(populate_by_name=True, protected_namespaces=())74 75 @model_validator(mode="after")76 def validate_environment(self) -> Self:77 """Validate that auth token exists in environment."""78 session = requests.Session()79 session.headers.update(80 {81 "Authorization": f"Bearer {self.baichuan_api_key.get_secret_value()}",82 "Accept-Encoding": "identity",83 "Content-type": "application/json",84 }85 )86 self.session = session87 return self88 89 def _embed(self, texts: List[str]) -> Optional[List[List[float]]]:90 """Internal method to call Baichuan Embedding API and return embeddings.91 92 Args:93 texts: A list of texts to embed.94 95 Returns:96 A list of list of floats representing the embeddings, or None if an97 error occurs.98 """99 chunk_texts = [100 texts[i : i + self.chunk_size]101 for i in range(0, len(texts), self.chunk_size)102 ]103 embed_results = []104 for chunk in chunk_texts:105 response = self.session.post(106 BAICHUAN_API_URL, json={"input": chunk, "model": self.model_name}107 )108 # Raise exception if response status code from 400 to 600109 response.raise_for_status()110 # Check if the response status code indicates success111 if response.status_code == 200:112 resp = response.json()113 embeddings = resp.get("data", [])114 # Sort resulting embeddings by index115 sorted_embeddings = sorted(embeddings, key=lambda e: e.get("index", 0))116 # Return just the embeddings117 embed_results.extend(118 [result.get("embedding", []) for result in sorted_embeddings]119 )120 else:121 # Log error or handle unsuccessful response appropriately122 # Handle 100 <= status_code < 400, not include 200123 raise RequestException(124 f"Error: Received status code {response.status_code} from "125 "`BaichuanEmbedding` API"126 )127 return embed_results128 129 def embed_documents(self, texts: List[str]) -> Optional[List[List[float]]]: # type: ignore[override]130 """Public method to get embeddings for a list of documents.131 132 Args:133 texts: The list of texts to embed.134 135 Returns:136 A list of embeddings, one for each text, or None if an error occurs.137 """138 return self._embed(texts)139 140 def embed_query(self, text: str) -> Optional[List[float]]: # type: ignore[override]141 """Public method to get embedding for a single query text.142 143 Args:144 text: The text to embed.145 146 Returns:147 Embeddings for the text, or None if an error occurs.148 """149 result = self._embed([text])150 return result[0] if result is not None else None151 