| 1358 | return new_norm_ema, new_square_ema |
| 1359 | |
| 1360 | def _is_valid_step( |
| 1361 | g_norm: Tensor, |
| 1362 | drop_norm: Union[float, DropNormThresholdFn], |
| 1363 | *, |
| 1364 | norm_ema: Optional[Tensor], |
| 1365 | norm_square_ema: Optional[Tensor], |
| 1366 | count: Optional[Tensor], |
| 1367 | drop_stats: Optional[dict[str, Tensor]], |
| 1368 | ) -> tuple[Tensor, Optional[dict[str, Tensor]]]: |
| 1369 | if isinstance(drop_norm, (float, int)): |
| 1370 | return g_norm < drop_norm, None |
| 1371 | else: |
| 1372 | stddev = _stddev(norm_ema, norm_square_ema) |
| 1373 | thresholds = drop_norm(count=count, mean=norm_ema, stddev=stddev) |
| 1374 | new_drop_stats = {} |
| 1375 | is_valid = None |
| 1376 | for key, val in thresholds.items(): |
| 1377 | less = g_norm < val |
| 1378 | is_valid = less if is_valid is None else jnp.logical_and(is_valid, less) |
| 1379 | new_drop_stats[key] = jnp.where( |
| 1380 | less, |
| 1381 | drop_stats[key], |
| 1382 | optax.safe_int32_increment(drop_stats[key]), |
| 1383 | ) |
| 1384 | return is_valid, new_drop_stats |
| 1385 | |
| 1386 | # Check if every gradient is finite. |
| 1387 | flat_updates = jax.tree_util.tree_flatten(updates)[0] |