MCPcopy Create free account
hub / github.com/OpenGVLab/HumanBench / _approx_sq_grad

Method _approx_sq_grad

PATH/core/optimizers/adafactor.py:117–129  ·  view source on GitHub ↗
(
        self,
        exp_avg_sq_row: torch.Tensor,
        exp_avg_sq_col: torch.Tensor,
        output: torch.Tensor,
    )

Source from the content-addressed store, hash-verified

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.

Callers 1

stepMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected