如何保存/加载Faiss KMeans模型以用于后续推理
Faiss KMeans模型无法存储与加载的解决方案
你遇到的报错是因为faiss.Kmeans训练后得到的index是内部专用索引类型,无法直接用faiss.write_index序列化。下面提供两种可行的解决方法:
方法一:手动保存聚类中心,重建KMeans模型
这种方法保留KMeans对象的完整结构,适合需要后续继续训练或使用KMeans内置方法的场景。
训练与保存代码
import numpy as np import faiss # 训练KMeans clustering = faiss.Kmeans(candles.shape[1], k=clusters, niter=epochs, gpu=gpu, verbose=True) clustering.train(X) # 保存聚类中心 np.save(f"{out_file}_centroids.npy", clustering.centroids) # 保存模型关键参数(维度、聚类数) with open(f"{out_file}_params.txt", "w") as f: f.write(f"{candles.shape[1]}\n{clusters}")
加载与推理代码
# 加载参数 with open(f"{out_file}_params.txt", "r") as f: dim = int(f.readline()) k = int(f.readline()) # 加载聚类中心 centroids = np.load(f"{out_file}_centroids.npy") # 重建KMeans模型 clustering_loaded = faiss.Kmeans(dim, k, niter=0, gpu=gpu, verbose=False) clustering_loaded.centroids = centroids # 构建用于搜索的索引 clustering_loaded.index = faiss.IndexFlatL2(dim) clustering_loaded.index.add(centroids) # 执行聚类搜索 distances, labels = clustering_loaded.index.search(x, 1)
方法二:转换为可序列化的标准索引
如果只需要用聚类中心做搜索推理,直接把聚类中心存入标准的IndexFlatL2索引即可,无需保留KMeans对象。
保存代码
# 创建可序列化的FlatL2索引 save_index = faiss.IndexFlatL2(candles.shape[1]) # 添加聚类中心到索引 save_index.add(clustering.centroids) # 保存索引 faiss.write_index(save_index, f"{out_file}.faiss")
加载与推理代码
# 加载索引 model2 = faiss.read_index(f"{out_file}.faiss") # 执行搜索 distances, labels = model2.search(x, 1)
注意事项
- 如果训练时使用了GPU,加载时需确保
gpu参数设置一致,避免设备不匹配问题。 - 方法一的
niter=0是为了避免加载后自动训练,若需要后续继续训练,可调整为对应迭代次数。
内容的提问来源于stack exchange,提问作者KIC
相关产品推荐
相关产品推荐

