codekingpro/portable-devtools
114k
1"""Wrapper around Xinference embedding models."""2 3from typing import Any, List, Optional4 5from langchain_core.embeddings import Embeddings6 7 8class XinferenceEmbeddings(Embeddings):9 """Xinference embedding models.10 11 To use, you should have the xinference library installed:12 13 .. code-block:: bash14 15 pip install xinference16 17 If you're simply using the services provided by Xinference, you can utilize the xinference_client package:18 19 .. code-block:: bash20 21 pip install xinference_client22 23 Check out: https://github.com/xorbitsai/inference24 To run, you need to start a Xinference supervisor on one server and Xinference workers on the other servers.25 26 Example:27 To start a local instance of Xinference, run28 29 .. code-block:: bash30 31 $ xinference32 33 You can also deploy Xinference in a distributed cluster. Here are the steps:34 35 Starting the supervisor:36 37 .. code-block:: bash38 39 $ xinference-supervisor40 41 If you're simply using the services provided by Xinference, you can utilize the xinference_client package:42 43 .. code-block:: bash44 45 pip install xinference_client46 47 Starting the worker:48 49 .. code-block:: bash50 51 $ xinference-worker52 53 Then, launch a model using command line interface (CLI).54 55 Example:56 57 .. code-block:: bash58 59 $ xinference launch -n orca -s 3 -q q4_060 61 It will return a model UID. Then you can use Xinference Embedding with LangChain.62 63 Example:64 65 .. code-block:: python66 67 from langchain_community.embeddings import XinferenceEmbeddings68 69 xinference = XinferenceEmbeddings(70 server_url="http://0.0.0.0:9997",71 model_uid = {model_uid} # replace model_uid with the model UID return from launching the model72 )73 74 """ # noqa: E50175 76 client: Any77 server_url: Optional[str]78 """URL of the xinference server"""79 model_uid: Optional[str]80 """UID of the launched model"""81 82 def __init__(83 self, server_url: Optional[str] = None, model_uid: Optional[str] = None84 ):85 try:86 from xinference.client import RESTfulClient87 except ImportError:88 try:89 from xinference_client import RESTfulClient90 except ImportError as e:91 raise ImportError(92 "Could not import RESTfulClient from xinference. Please install it"93 " with `pip install xinference` or `pip install xinference_client`."94 ) from e95 96 super().__init__()97 98 if server_url is None:99 raise ValueError("Please provide server URL")100 101 if model_uid is None:102 raise ValueError("Please provide the model UID")103 104 self.server_url = server_url105 106 self.model_uid = model_uid107 108 self.client = RESTfulClient(server_url)109 110 def embed_documents(self, texts: List[str]) -> List[List[float]]:111 """Embed a list of documents using Xinference.112 Args:113 texts: The list of texts to embed.114 Returns:115 List of embeddings, one for each text.116 """117 118 model = self.client.get_model(self.model_uid)119 120 embeddings = [121 model.create_embedding(text)["data"][0]["embedding"] for text in texts122 ]123 return [list(map(float, e)) for e in embeddings]124 125 def embed_query(self, text: str) -> List[float]:126 """Embed a query of documents using Xinference.127 Args:128 text: The text to embed.129 Returns:130 Embeddings for the text.131 """132 133 model = self.client.get_model(self.model_uid)134 135 embedding_res = model.create_embedding(text)136 137 embedding = embedding_res["data"][0]["embedding"]138 139 return list(map(float, embedding))140 