You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.30 07:17:27