(grad, name)
| 82 | ) |
| 83 | |
| 84 | def detect_nan_hook(grad, name): |
| 85 | if torch.isnan(grad).any(): |
| 86 | print(f"NaN detected in gradients of {name}!") |
| 87 | print(f"Gradient values: {grad}") |
| 88 | exit() |
| 89 | |
| 90 | # 注册钩子到每个参数 |
| 91 | for name, param in model.named_parameters(): |