修改LM Studio向量错误问题
This commit is contained in:
+48
-7
@@ -120,18 +120,59 @@ class OllamaEmbeddingAdapter(BaseEmbeddingAdapter):
|
|||||||
return []
|
return []
|
||||||
|
|
||||||
class MLStudioEmbeddingAdapter(BaseEmbeddingAdapter):
|
class MLStudioEmbeddingAdapter(BaseEmbeddingAdapter):
|
||||||
|
"""
|
||||||
|
基于 LM Studio 的 embedding 适配器
|
||||||
|
"""
|
||||||
def __init__(self, api_key: str, base_url: str, model_name: str):
|
def __init__(self, api_key: str, base_url: str, model_name: str):
|
||||||
self._embedding = OpenAIEmbeddings(
|
self.url = ensure_openai_base_url_has_v1(base_url)
|
||||||
openai_api_key=api_key,
|
if not self.url.endswith('/embeddings'):
|
||||||
openai_api_base=ensure_openai_base_url_has_v1(base_url),
|
self.url = f"{self.url}/embeddings"
|
||||||
model=model_name
|
|
||||||
)
|
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]]:
|
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]:
|
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):
|
class GeminiEmbeddingAdapter(BaseEmbeddingAdapter):
|
||||||
"""
|
"""
|
||||||
|
|||||||
Reference in New Issue
Block a user