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

Function ema

axlearn/common/optimizers.py:888–1055  ·  view source on GitHub ↗

Compute an exponential moving average of updates. Based on: Args: decay: the decay rate schedule for th

(
    decay: schedule.Schedule,
    debias: bool = True,
    accumulator_dtype: Optional[jnp.dtype] = jnp.float32,
)

Source from the content-addressed store, hash-verified

886
887
888def ema(
889 decay: schedule.Schedule,
890 debias: bool = True,
891 accumulator_dtype: Optional[jnp.dtype] = jnp.float32,
892) -> PartitionedGradientTransformation:
893 """Compute an exponential moving average of updates.
894
895 Based on:
896 <https://github.com/deepmind/optax/blob/252d15/optax/_src/transform.py#L120-L158>
897 <https://github.com/tensorflow/lingvo/blob/4c0252f/lingvo/jax/optimizers.py#L1780-L1858>
898
899 Args:
900 decay: the decay rate schedule for the exponential moving average.
901 debias: whether to debias the transformed gradient.
902 accumulator_dtype: optional `dtype` to use for the accumulator; if `None`
903 or if the parameter is a scalar then the `dtype` is inferred.
904 Supports: float32, bfloat16, int16, int8.
905
906 Returns:
907 A corresponding `PartitionedGradientTransformation`.
908
909 Raises:
910 ValueError: If accumulator_dtype is invalid.
911 """
912 decay_fn = schedule.as_schedule_fn(decay)
913 # Validate accumulator_dtype.
914 float_dtypes = [jnp.float32, jnp.bfloat16]
915 int_dtypes = [jnp.int16, jnp.int8]
916 if accumulator_dtype is not None:
917 valid_dtypes = float_dtypes + int_dtypes
918 accumulator_dtype = jax.dtypes.canonicalize_dtype(accumulator_dtype)
919 if accumulator_dtype not in valid_dtypes:
920 raise ValueError(f"accumulator_dtype must be one of {valid_dtypes} if set.")
921
922 def _should_quantize(t_shape: Sequence[int]):
923 return t_shape and accumulator_dtype in int_dtypes
924
925 @dataclasses.dataclass
926 class _TensorEma:
927 # The exponential moving average state and quantization scaling factor for a tensor.
928 value: Tensor # Current value of the momentum estimate for the tensor.
929 qstep_size: Tensor # Scaling factor, for converting 'value' to float if quantized, else 0.
930
931 def _to_state(count: Tensor, ema_tree: NestedTree):
932 return EmaState(
933 count=count,
934 ema=jax.tree.map(lambda ema: ema.value, ema_tree),
935 scale=jax.tree.map(lambda ema: ema.qstep_size, ema_tree),
936 )
937
938 def init_fn(params):
939 def _init(t):
940 # Store momentum in accumulator_dtype if it is set and p is not scalar.
941 if t.shape and accumulator_dtype is not None:
942 value = jnp.zeros(t.shape, dtype=accumulator_dtype)
943 else:
944 value = jnp.zeros(t.shape, dtype=t.dtype)
945

Callers 2

test_ema_parityMethod · 0.90
adafactor_optimizerFunction · 0.85

Tested by 1

test_ema_parityMethod · 0.72