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

使用simpletransformers时如何在wandb中记录各阶段的模型工件?

基于simpletransformers + wandb 记录QA模型全阶段工件的实现方法

1. 前置准备

确保已经安装对应依赖,提前完成wandb本地身份验证:

pip install simpletransformers wandb
wandb login 你的wandb账户API密钥

2. 训练开始前记录数据集工件

在初始化模型之前先初始化wandb运行实例,上传训练/验证/测试数据集:

import wandb
import os
from simpletransformers.question_answering import QuestionAnsweringModel

# 初始化wandb运行
run = wandb.init(
    project="你的项目名",
    name="本次运行的名称,比如qa-model-bert-base-0618",
    job_type="train"
)

# 记录数据集工件
dataset_artifact = wandb.Artifact(
    name="qa-datasets",
    type="dataset",
    description="问答模型训练、验证、测试数据集"
)
# 依次添加三个数据集文件
dataset_artifact.add_file("train.json")
dataset_artifact.add_file("eval.json")
dataset_artifact.add_file("test.json")
run.log_artifact(dataset_artifact)

3. 配置simpletransformers的wandb集成参数

定义QA模型训练参数,关联当前已经初始化的wandb运行:

model_args = {
    "output_dir": "output/", # 输出文件保存目录,后续要从这个目录读预测结果和最优模型
    "best_model_dir": "output/best_model/", # 最优模型保存目录
    # 其他常规训练参数,比如学习率、批次大小、训练轮次等按需配置
    "num_train_epochs": 5,
    "train_batch_size": 16,
    "evaluate_during_training": True,
    # wandb关联配置
    "wandb_project": "你的项目名",
    "wandb_kwargs": {"run": run} # 复用前面初始化的run,避免重复创建运行实例
}

# 初始化QA模型,按需选择模型类型和预训练权重
model = QuestionAnsweringModel(
    "bert",
    "bert-base-cased",
    args=model_args,
    use_cuda=True # 没有GPU可以设为False
)

4. 训练完成后记录预测结果和最优模型

先完成训练、测试流程,之后把生成的预测文件和最优模型上传为wandb工件:

# 训练模型
model.train_model(train_data, eval_data=eval_data)

# 测试集推理,生成预测结果
result, nbest_preds, _ = model.eval_model(test_data)

# 记录测试集预测结果工件
pred_artifact = wandb.Artifact(
    name="qa-test-predictions",
    type="prediction",
    description="测试集n_best预测结果"
)
pred_artifact.add_file("output/nbest_predictions_test.json")
run.log_artifact(pred_artifact)

# 记录最优模型工件
model_artifact = wandb.Artifact(
    name="qa-best-model",
    type="model",
    description="验证集表现最优的问答模型"
)
# 把整个最优模型目录添加到工件
model_artifact.add_dir("output/best_model/")
run.log_artifact(model_artifact)

# 结束wandb运行
run.finish()

注意事项

  • 所有文件路径需要和你本地实际存储路径一致,如果文件不在代码运行的当前目录,需要写完整绝对路径
  • 可以给工件添加别名,比如创建Artifact时传入aliases=["best", "v1.0"],后续可以直接通过别名拉取对应版本的工件
  • 如果需要记录中间训练过程生成的其他工件,按照上面的逻辑创建对应Artifact并上传即可

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 00:15:04