使用MLflow保存含Universal Sentence Encoder的Top2Vec模型遇Pickle错误
问题解决方案
问题根源
当Top2Vec使用universal-sentence-encoder作为嵌入模型时,内部会持有TensorFlow的模型实例,而TensorFlow的重复消息字段无法被Pickle序列化,这就是MLflow保存时抛出PickleError的原因。不使用USE时,Top2Vec的模型结构可正常被Pickle处理,因此保存操作无异常。
解决步骤
核心思路是利用Top2Vec自带的save/load方法处理USE模型的序列化,避免直接Pickle整个模型实例,同时通过MLflow的artifact机制托管Top2Vec模型文件。
1. 训练并保存Top2Vec核心模型
训练完成后,先将Top2Vec模型保存到本地目录,该方法会自动处理USE模型的序列化(保存为TensorFlow SavedModel格式):
from top2vec import Top2Vec import mlflow.pyfunc import os # 训练Top2Vec模型(带USE嵌入) model = Top2Vec(documents=doc_list, document_ids=ids_list, embedding_model='universal-sentence-encoder') # 保存Top2Vec模型到本地临时目录 top2vec_local_path = "./top2vec_saved_model" model.save(top2vec_local_path)
2. 修改MLflow PythonModel类
重写load_context方法,在模型加载时通过Top2Vec的load方法重新初始化模型,而非直接持有训练好的实例:
max_num_recomm = 50 class Top2VecModel(mlflow.pyfunc.PythonModel): def load_context(self, context): # 从MLflow artifact加载Top2Vec模型 from top2vec import Top2Vec self._tm = Top2Vec.load(context.artifacts["top2vec_model"]) def d_by_d_search(self, id): results = self._tm.search_documents_by_documents(doc_ids=id, num_docs=max_num_recomm) return results def d_by_kw_search(self, kw): results = self._tm.search_documents_by_keywords(keywords=kw, num_docs=max_num_recomm) return results
3. 保存MLflow模型
将本地保存的Top2Vec模型作为artifact传入MLflow的保存方法:
mlflow.pyfunc.save_model( path='./mlflow_top2vec_final', python_model=Top2VecModel(), artifacts={"top2vec_model": top2vec_local_path} )
关键说明
- Top2Vec的
save/load方法会专门处理USE等TensorFlow嵌入模型的序列化逻辑,避免Pickle直接处理TF对象导致的错误。 - 你之前尝试保存
model.embed()的方式错误,因为embed()是生成嵌入的方法,而非模型文件路径,因此会触发输入签名不匹配的异常。
内容的提问来源于stack exchange,提问作者Javid Jouzdani
相关产品推荐
相关产品推荐

