如何在PyTorch+Transformers的run_glue.py BERT训练代码中接入Weights & Biases
Weights & Biases 自定义PyTorch训练循环配置步骤
- 第一步:安装并完成身份验证
先执行安装命令:pip install wandb
在代码开头导入库,调用登录接口,按提示输入你WandB账号的API密钥即可:import wandb wandb.login() - 第二步:初始化训练运行,同步超参数
在模型、超参数定义完成后,初始化wandb运行实例,把所有训练相关的超参数传入config字段:run = wandb.init( project="你自定义的项目名称,比如电商SEO-BERTimbau训练", config={ "learning_rate": 2e-5, "train_batch_size": 32, "eval_batch_size": 64, "num_train_epochs": 10, "weight_decay": 1e-4, "model_name": "BERTimbau-base", # 其余你用到的所有超参数都可以补充到这里 } ) # 如果你是用argparse加载run_glue.py的原有参数,直接转字典传入即可,不用逐个写: # run = wandb.init(project="你的项目名", config=vars(args)) - 第三步:绑定模型,自动跟踪梯度与参数变化
模型实例化完成后,添加watch方法即可自动跟踪模型的梯度、权重变化:# model是你实例化的BERTimbau模型对象 wandb.watch(model, log="all", log_freq=100) # log_freq参数可自行调整,代表每多少个训练步同步一次参数 - 第四步:在训练、验证流程中添加指标上报
在自定义的训练循环中,每步训练结束后上报训练损失等指标:
每轮验证结束后,上报验证集的所有评估指标:# 原有训练步逻辑 loss = model(**batch)[0] loss.backward() optimizer.step() lr_scheduler.step() optimizer.zero_grad() global_step += 1 # 新增wandb指标上报逻辑 wandb.log({"train/loss": loss.item()}, step=global_step)# 原有验证逻辑执行完成,得到评估指标后 wandb.log({ "eval/loss": eval_avg_loss, "eval/accuracy": eval_acc, "eval/f1": eval_f1, # 其余你需要跟踪的验证指标都可以补充在这里 }, step=global_step) - 可选配置:同步模型 checkpoint 到WandB
如果你需要把训练中保存的模型文件同步到WandB后台,保存checkpoint后调用save方法即可:torch.save(model.state_dict(), "./bertimbau_seo_epoch3.pt") wandb.save("./bertimbau_seo_epoch3.pt")
如果GCP虚拟机无法访问公网,可在初始化wandb时添加mode="offline"参数开启离线模式,训练结束后在虚拟机终端执行wandb sync命令即可批量同步所有训练日志。
内容的提问来源于stack exchange,提问作者Guilherme Giuliano Nicolau
相关产品推荐
相关产品推荐

