fix: change embedding to embedding_function in FAISS

This commit is contained in:
Gabriel Luiz Freitas Almeida 2023-11-27 18:02:32 -03:00
commit 45556f9b5a

View file

@ -1,17 +1,17 @@
from typing import Any, Callable, Dict, Type
from langchain.vectorstores import (
Pinecone,
Qdrant,
Chroma,
FAISS,
Weaviate,
SupabaseVectorStore,
MongoDBAtlasVectorSearch,
)
from langchain.schema import Document
import os import os
from typing import Any, Callable, Dict, Type
import orjson import orjson
from langchain.schema import Document
from langchain.vectorstores import (
FAISS,
Chroma,
MongoDBAtlasVectorSearch,
Pinecone,
Qdrant,
SupabaseVectorStore,
Weaviate,
)
def docs_in_params(params: dict) -> bool: def docs_in_params(params: dict) -> bool:
@ -28,8 +28,8 @@ def initialize_mongodb(class_object: Type[MongoDBAtlasVectorSearch], params: dic
MONGODB_ATLAS_CLUSTER_URI = params.pop("mongodb_atlas_cluster_uri") MONGODB_ATLAS_CLUSTER_URI = params.pop("mongodb_atlas_cluster_uri")
if not MONGODB_ATLAS_CLUSTER_URI: if not MONGODB_ATLAS_CLUSTER_URI:
raise ValueError("Mongodb atlas cluster uri must be provided in the params") raise ValueError("Mongodb atlas cluster uri must be provided in the params")
from pymongo import MongoClient
import certifi import certifi
from pymongo import MongoClient
client: MongoClient = MongoClient( client: MongoClient = MongoClient(
MONGODB_ATLAS_CLUSTER_URI, tlsCAFile=certifi.where() MONGODB_ATLAS_CLUSTER_URI, tlsCAFile=certifi.where()
@ -120,6 +120,7 @@ def initialize_faiss(class_object: Type[FAISS], params: dict):
return class_object.load_local return class_object.load_local
save_local = params.get("save_local") save_local = params.get("save_local")
params["embedding_function"] = params.pop("embedding")
faiss_index = class_object(**params) faiss_index = class_object(**params)
if save_local: if save_local:
faiss_index.save_local(folder_path=save_local) faiss_index.save_local(folder_path=save_local)