MCPcopy Create free account
hub / github.com/apple/axlearn / _is_valid_step

Function _is_valid_step

axlearn/common/optimizers.py:1360–1384  ·  view source on GitHub ↗
(
            g_norm: Tensor,
            drop_norm: Union[float, DropNormThresholdFn],
            *,
            norm_ema: Optional[Tensor],
            norm_square_ema: Optional[Tensor],
            count: Optional[Tensor],
            drop_stats: Optional[dict[str, Tensor]],
        )

Source from the content-addressed store, hash-verified

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]

Callers 1

update_fnFunction · 0.85

Calls 2

_stddevFunction · 0.85
itemsMethod · 0.80

Tested by

no test coverage detected