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

如何用MLflow记录自定义PyTorch模型?PreSumm模型日志问题咨询

1. 如何使用MLflow记录自定义PyTorch模型?

其实步骤挺清晰的,咱一步步来:

  • 先把依赖装齐:pip install mlflow torch,确保MLflow和PyTorch都在你的环境里。
  • 开启MLflow运行:更推荐用上下文管理器,这样不用手动结束运行,省心不少:
    with mlflow.start_run(run_name="我的自定义模型训练"):
        # 所有记录操作都放在这个上下文里
    
  • 准备好你的自定义模型——记住必须是torch.nn.Module的子类哦,比如写个简单的线性模型示例:
    class CustomTorchModel(torch.nn.Module):
        def __init__(self, input_dim):
            super().__init__()
            self.fc = torch.nn.Linear(input_dim, 1)
        def forward(self, x):
            return self.fc(x)
    
    model = CustomTorchModel(input_dim=10)
    
  • 训练完模型后,直接调用mlflow.pytorch.log_model(model, "保存的模型目录名"),就能把模型记录到MLflow的实验中了。
  • 额外加分项:你还可以顺便记录训练参数、指标,比如mlflow.log_param("学习率", 0.001),mlflow.log_metric("验证集准确率", 0.92),这样后续复盘实验时能更直观对比。
2. 解决PreSumm中trainer非Module类无法用mlflow.pytorch.log_model的问题

这个问题我之前也碰到过,核心思路其实很简单:mlflow.pytorch.log_model只认torch.nn.Module的子类实例,所以你直接传AbsSummarizer模型就行,不用管trainer。

具体操作可以这么做:

  • 先找到你代码里初始化AbsSummarizer的地方,在train_abstractive.py里应该有类似这样的代码:
    model = AbsSummarizer(args, device, load_pretrained_emb)
    
    这个model就是正经的torch.nn.Module子类,也是trainer内部实际用来训练的核心模型,训练过程中它的权重会被正常更新。
  • 训练结束后,直接把这个model传给MLflow的log_model方法就行,比如:
    import mlflow.pytorch
    
    with mlflow.start_run(run_name="PreSumm摘要训练"):
        # 先记录点训练参数和指标,比如batch_size、验证集ROUGE分数
        mlflow.log_param("batch_size", args.batch_size)
        mlflow.log_metric("val_rouge1", 你的验证集ROUGE1分数)
        # 关键一步:记录AbsSummarizer模型
        mlflow.pytorch.log_model(model, "abstractive_summarizer")
    
  • 后续要加载模型推理的话,用mlflow.pytorch.load_model("模型的MLflow路径"),加载出来的模型就能直接用它的推理方法(比如generate),和你平时用PyTorch模型没啥区别。

要是你担心trainer里的一些配置或者优化器状态要保存,也可以用mlflow.log_artifact()把这些文件单独归档,比如保存优化器状态:

# 先把优化器状态存成文件
torch.save(trainer.optimizer.state_dict(), "optimizer_state.pth")
# 再把这个文件作为artifact传到MLflow
mlflow.log_artifact("optimizer_state.pth")

内容的提问来源于stack exchange,提问作者Chaine

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.07 21:03:15