使用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
相关产品推荐
相关产品推荐

