wandb.watch() s’attache aux paramètres et aux gradients d’un modèle PyTorch et journalise à intervalles réguliers des histogrammes de leurs valeurs. Cette fonction permet de diagnostiquer l’instabilité de l’entraînement, la disparition des gradients et les neurones morts.
Utilisation de base
Appelez wandb.watch() après wandb.init() et avant la première étape d’entraînement :
log_freq lots (la valeur par défaut de Run.watch() est log_freq=1000 ; l’exemple utilise 100 pour obtenir un retour plus rapide). Ils apparaissent dans l’onglet Charts sous des clés telles que gradients/layer_name.weight.
Options du paramètre log
log_graph=True si vous souhaitez obtenir le graphe de calcul alors que la journalisation des histogrammes est désactivée ou minimale. Le graphe s’affiche dans l’onglet Overview du run, sous Model. Consultez Run.watch() pour savoir comment log, log_graph et log_freq interagissent.
wandb.watch() séparément pour chaque modèle (utile pour l’entraînement de GAN) :
log_freq. Journaliser à chaque étape (log_freq=1) peut ralentir considérablement l’entraînement. Pour la plupart des runs d’entraînement, une valeur comprise entre 50 et 200 est habituelle. Si les performances sont critiques, définissez log="parameters" plutôt que log="gradients" : les histogrammes de paramètres sont calculés sans hook sur la rétropropagation et sont donc moins coûteux.
Arrêter la surveillance
Pour arrêter la journalisation des gradients en cours d’entraînement :
Experiments Métriques Runs