(x: torch.Tensor, eps: float =1e-20)
| 26 | return 2*dot(x, n)*n - x |
| 27 | |
| 28 | def length(x: torch.Tensor, eps: float =1e-20) -> torch.Tensor: |
| 29 | return torch.sqrt(torch.clamp(dot(x,x), min=eps)) # Clamp to avoid nan gradients because grad(sqrt(0)) = NaN |
| 30 | |
| 31 | def safe_normalize(x: torch.Tensor, eps: float =1e-20) -> torch.Tensor: |
| 32 | return x / length(x, eps) |