From 140a86206d70edee489e3d5b6f6d1f7cda870840 Mon Sep 17 00:00:00 2001 From: CNlaojing Date: Tue, 25 Mar 2025 23:47:23 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E6=94=B9LM=20Studio=E5=90=91=E9=87=8F?= =?UTF-8?q?=E9=94=99=E8=AF=AF=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- embedding_adapters.py | 55 +++++++++++++++++++++++++++++++++++++------ 1 file changed, 48 insertions(+), 7 deletions(-) diff --git a/embedding_adapters.py b/embedding_adapters.py index 2294de9..869dc1d 100644 --- a/embedding_adapters.py +++ b/embedding_adapters.py @@ -120,18 +120,59 @@ class OllamaEmbeddingAdapter(BaseEmbeddingAdapter): return [] class MLStudioEmbeddingAdapter(BaseEmbeddingAdapter): + """ + 基于 LM Studio 的 embedding 适配器 + """ def __init__(self, api_key: str, base_url: str, model_name: str): - self._embedding = OpenAIEmbeddings( - openai_api_key=api_key, - openai_api_base=ensure_openai_base_url_has_v1(base_url), - model=model_name - ) + self.url = ensure_openai_base_url_has_v1(base_url) + if not self.url.endswith('/embeddings'): + self.url = f"{self.url}/embeddings" + + self.headers = { + "Authorization": f"Bearer {api_key}", + "Content-Type": "application/json" + } + self.model_name = model_name def embed_documents(self, texts: List[str]) -> List[List[float]]: - return self._embedding.embed_documents(texts) + try: + payload = { + "input": texts, + "model": self.model_name + } + response = requests.post(self.url, json=payload, headers=self.headers) + response.raise_for_status() + result = response.json() + if "data" not in result: + logging.error(f"Invalid response format from LM Studio API: {result}") + return [[]] * len(texts) + return [item.get("embedding", []) for item in result["data"]] + except requests.exceptions.RequestException as e: + logging.error(f"LM Studio API request failed: {str(e)}") + return [[]] * len(texts) + except (KeyError, IndexError, ValueError, TypeError) as e: + logging.error(f"Error parsing LM Studio API response: {str(e)}") + return [[]] * len(texts) def embed_query(self, query: str) -> List[float]: - return self._embedding.embed_query(query) + try: + payload = { + "input": query, + "model": self.model_name + } + response = requests.post(self.url, json=payload, headers=self.headers) + response.raise_for_status() + result = response.json() + if "data" not in result or not result["data"]: + logging.error(f"Invalid response format from LM Studio API: {result}") + return [] + return result["data"][0].get("embedding", []) + except requests.exceptions.RequestException as e: + logging.error(f"LM Studio API request failed: {str(e)}") + return [] + except (KeyError, IndexError, ValueError, TypeError) as e: + logging.error(f"Error parsing LM Studio API response: {str(e)}") + return [] class GeminiEmbeddingAdapter(BaseEmbeddingAdapter): """