应该何时调用wandb.watch以正确跟踪模型参数与梯度?
问题根因
你调用wandb.watch()不生效主要有两个核心原因:
- 缺少
wandb.init()初始化步骤:wandb.watch()必须绑定到一个已经初始化完成的WandB运行实例才能工作,你的代码全程没有调用wandb.init(),运行实例未创建,自然无法追踪模型参数和梯度。 - 默认记录频率不匹配你的测试场景:
wandb.watch()默认的参数/梯度记录频率为log_freq=1000,你的测试训练步数只有12步,远低于默认频率阈值,就算完成初始化也不会触发记录。 - 分布式场景注意事项:如果是分布式训练场景,
wandb.watch()和所有WandB相关操作都只能在主进程调用,子进程调用不会生效还可能引发冲突。
修复方案
你只需要调整debug_test函数的逻辑,补充初始化代码、调整watch参数即可:
def debug_test(): args: Namespace = get_args() args.num_its = 12 # 新增:主进程初始化WandB运行实例,必须在watch之前执行 if is_lead_worker(args.rank): wandb.init( project='playground', entity='brando', config=vars(args) # 自动同步超参数 ) # - get mdl, opt, scheduler, etc args.mdl = get_simple_model(in_features=5, hidden_features=20, out_features=1, num_layer=2) # 调整watch参数:log='all'同时记录参数和梯度,log_freq=1适配少步数测试 if is_lead_worker(args.rank): wandb.watch(args.mdl, log='all', log_freq=1) args.optimizer = torch.optim.Adam(args.mdl.parameters(), lr=1e-1) args.scheduler = torch.optim.lr_scheduler.ExponentialLR(args.optimizer, gamma=0.999, verbose=False) # 后续训练逻辑保持不变 ...
验证方法
修复后运行代码,进入对应WandB运行的详情页,切换到模型标签页,即可查看各层参数、梯度的变化曲线。
内容的提问来源于stack exchange,提问作者Charlie Parker
相关产品推荐
相关产品推荐

