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,
)
| 886 | |
| 887 | |
| 888 | def 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 |