| 115 | return tensor.norm(2) / (tensor.numel() ** 0.5) |
| 116 | |
| 117 | def _approx_sq_grad( |
| 118 | self, |
| 119 | exp_avg_sq_row: torch.Tensor, |
| 120 | exp_avg_sq_col: torch.Tensor, |
| 121 | output: torch.Tensor, |
| 122 | ) -> None: |
| 123 | r_factor = ( |
| 124 | (exp_avg_sq_row / exp_avg_sq_row.mean(dim=-1, keepdim=True)) |
| 125 | .rsqrt_() |
| 126 | .unsqueeze(-1) |
| 127 | ) |
| 128 | c_factor = exp_avg_sq_col.unsqueeze(-2).rsqrt() |
| 129 | torch.mul(r_factor, c_factor, out=output) |
| 130 | |
| 131 | def step(self, closure: OptLossClosure = None) -> OptFloat: |
| 132 | r"""Performs a single optimization step. |