Computes raw updates from gradients.
(grad: Tensor, pps: _AdastarPerParamState)
| 1920 | return updates, pps_tree |
| 1921 | |
| 1922 | def _raw_updates(grad: Tensor, pps: _AdastarPerParamState) -> _AdastarUpdateResult: |
| 1923 | """Computes raw updates from gradients.""" |
| 1924 | smoothed_gradient, gradient_ema = _moment( |
| 1925 | grad, |
| 1926 | acc=pps.gradient_ema, |
| 1927 | decay=gradient_ema_decay, |
| 1928 | debias=gradient_ema_debias, |
| 1929 | ) |
| 1930 | smoothed_gradient_square, gradient_square_ema = _moment( |
| 1931 | grad**2 + eps_square, |
| 1932 | acc=pps.gradient_square_ema, |
| 1933 | decay=gradient_square_ema_decay, |
| 1934 | debias=gradient_square_ema_debias, |
| 1935 | ) |
| 1936 | raw_updates = smoothed_gradient / ((smoothed_gradient_square) ** 0.5 + eps) |
| 1937 | if logging.vlog_is_on(3): |
| 1938 | jax.debug.print("adastar mu={mu} nu={nu}", mu=gradient_ema, nu=gradient_square_ema) |
| 1939 | jax.debug.print("adastar raw_updates={u}", u=raw_updates) |
| 1940 | new_pps = _AdastarPerParamState( |
| 1941 | gradient_ema=gradient_ema, |
| 1942 | gradient_square_ema=gradient_square_ema, |
| 1943 | update_ema=pps.update_ema, |
| 1944 | ) |
| 1945 | return _AdastarUpdateResult(updates=raw_updates, pps=new_pps) |
| 1946 | |
| 1947 | def _smoothed_updates( |
| 1948 | raw_updates: Tensor, pps: _AdastarPerParamState |
no test coverage detected