wandb.watch() は PyTorch モデルのパラメーターと勾配にフックを追加し、それらの値のヒストグラムを一定間隔でログします。これは、トレーニングの不安定性、勾配消失、活動しないニューロンの診断に役立ちます。
基本的な使用方法
wandb.init() の後、最初のトレーニングステップの前に wandb.watch() を呼び出します:
log_freq バッチごとにログされます (Run.watch() のデフォルトは log_freq=1000 です。この例では、より早くフィードバックを得るために 100 を使用しています) 。これらは Charts タブに、gradients/layer_name.weight などのキーで表示されます。
log パラメーターのオプション
log_graph=True を渡します。グラフは run の Overview タブの Model で確認できます。log、log_graph、log_freq の相互作用については、Run.watch() を参照してください。
wandb.watch() を個別に呼び出します (GAN のトレーニングで便利です):
log_freq に比例するオーバーヘッドが発生します。毎ステップのログ (log_freq=1) は、トレーニングを大幅に遅くする可能性があります。ほとんどのトレーニング run では、50〜200 の値が一般的です。パフォーマンスが重要な場合は、log="gradients" ではなく log="parameters" を設定してください。パラメーターのヒストグラムは逆伝播のフックなしで計算されるため、負荷が低くなります。
監視の停止
トレーニングの途中で勾配のログを停止するには:
Experiments メトリクス Runs