(parameters, norm_type=2)
| 63 | |
| 64 | |
| 65 | def get_grad_norm(parameters, norm_type=2): |
| 66 | if isinstance(parameters, torch.Tensor): |
| 67 | parameters = [parameters] |
| 68 | parameters = list(filter(lambda p: p.grad is not None, parameters)) |
| 69 | norm_type = float(norm_type) |
| 70 | total_norm = 0 |
| 71 | for p in parameters: |
| 72 | param_norm = p.grad.data.norm(norm_type) |
| 73 | total_norm += param_norm.item() ** norm_type |
| 74 | total_norm = total_norm ** (1. / norm_type) |
| 75 | return total_norm |
| 76 | |
| 77 | |
| 78 | def auto_resume_helper(output_dir): |
no outgoing calls
no test coverage detected