如何高效存储与加载TfidfVectorizer生成的Scipy稀疏CSR矩阵?
解决Tfidf稀疏矩阵存储MongoDB及复用问题
核心需求
将TfidfVectorizer.fit_transform()生成的scipy.sparse.csr.csr_matrix存入MongoDB,后续可加载复用,用于新数据的transform及余弦相似度计算,同时保证内存效率。
方案一:序列化稀疏矩阵存入MongoDB
CSR稀疏矩阵无法直接存入MongoDB,可提取其核心结构数据(indices、indptr、data数组及矩阵shape),转换为MongoDB支持的格式存储,加载时再重建矩阵。
保存流程
import numpy as np from pymongo import MongoClient from sklearn.feature_extraction.text import TfidfVectorizer import pickle # 初始化并拟合数据 fit_data = ['this is some corpus data', 'foo', 'bar'] tf = TfidfVectorizer() tf_fitted = tf.fit_transform(fit_data) # 将CSR矩阵转换为可存储的结构 matrix_data = { "indices": tf_fitted.indices.tolist(), "indptr": tf_fitted.indptr.tolist(), "data": tf_fitted.data.tolist(), "shape": tf_fitted.shape, # 同时保存TfidfVectorizer实例,用于后续新数据转换 "vectorizer": pickle.dumps(tf) } # 存入MongoDB client = MongoClient("mongodb://localhost:27017/") db = client["tfidf_db"] collection = db["sparse_matrices"] collection.insert_one(matrix_data)
加载流程
from scipy.sparse import csr_matrix from pymongo import MongoClient import pickle from sklearn.metrics.pairwise import cosine_similarity # 从MongoDB读取数据 client = MongoClient("mongodb://localhost:27017/") db = client["tfidf_db"] collection = db["sparse_matrices"] saved_data = collection.find_one() # 重建CSR稀疏矩阵 loaded_matrix = csr_matrix( (np.array(saved_data["data"]), np.array(saved_data["indices"]), np.array(saved_data["indptr"])), shape=saved_data["shape"] ) # 加载TfidfVectorizer实例处理新数据 loaded_tf = pickle.loads(saved_data["vectorizer"]) transform_data = ['this is some other data to transform'] tf_transformed = loaded_tf.transform(transform_data) # 计算余弦相似度 similarity = cosine_similarity(loaded_matrix, tf_transformed).flatten() print(similarity)
方案二:直接序列化TfidfVectorizer实例(更推荐)
如果核心需求是减少重复拟合/向量化时间,无需存储拟合后的稀疏矩阵,直接序列化TfidfVectorizer实例更高效——只需保存实例的二进制数据,加载后可直接调用transform()处理新数据,也可随时重新生成拟合后的矩阵。
保存流程
import pickle from pymongo import MongoClient from sklearn.feature_extraction.text import TfidfVectorizer fit_data = ['this is some corpus data', 'foo', 'bar'] tf = TfidfVectorizer() # 拟合语料,让vectorizer学习特征 tf.fit(fit_data) # 序列化vectorizer为二进制 tf_pickle = pickle.dumps(tf) # 存入MongoDB client = MongoClient("mongodb://localhost:27017/") db = client["tfidf_db"] collection = db["vectorizers"] collection.insert_one({"vectorizer": tf_pickle})
加载流程
import pickle from pymongo import MongoClient from sklearn.feature_extraction.text import TfidfVectorizer from sklearn.metrics.pairwise import cosine_similarity # 从MongoDB读取序列化的vectorizer client = MongoClient("mongodb://localhost:27017/") db = client["tfidf_db"] collection = db["vectorizers"] saved_doc = collection.find_one() # 反序列化得到vectorizer实例 loaded_tf = pickle.loads(saved_doc["vectorizer"]) # 处理新数据 transform_data = ['this is some other data to transform'] tf_transformed = loaded_tf.transform(transform_data) # 如需原拟合矩阵,直接用实例生成(比重新拟合语料快) tf_fitted = loaded_tf.transform(fit_data) similarity = cosine_similarity(tf_fitted, tf_transformed).flatten() print(similarity)
注意事项
- 使用pickle序列化时,需确保前后环境的scikit-learn、numpy版本一致,避免反序列化失败。
- MongoDB存储二进制数据时会自动用
BinData类型处理,内存效率较高。
内容的提问来源于stack exchange,提问作者GUNTER
相关产品推荐
相关产品推荐

