def gradient_norm(module: nn.Module) -> float: parameter_norms = [ parameter.grad.detach().norm(2) for parameter in module.parameters() if parameter.grad is not None ] return torch.linalg.vector_norm(torch.stack(parameter_norms), 2).item()