| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| """K-FAC for accumulating statistics.""" |
| from typing import Any, Callable, Generic |
|
|
| import jax |
| import jax.numpy as jnp |
|
|
| from kfac_jax._src.utils import misc |
| from kfac_jax._src.utils import parallel |
| from kfac_jax._src.utils import types |
|
|
| Array = types.Array |
| Numeric = types.Numeric |
| Shape = types.Shape |
| DType = types.DType |
| ArrayTree = types.ArrayTree |
| TArrayTree = types.TArrayTree |
|
|
| AddFunction = Callable[[TArrayTree, TArrayTree, Numeric, Numeric], TArrayTree] |
|
|
|
|
| def default_add_function( |
| obj1: TArrayTree, |
| obj2: TArrayTree, |
| coeff1: Numeric, |
| coeff2: Numeric |
| ) -> TArrayTree: |
|
|
| return jax.tree_util.tree_map( |
| lambda x, y: coeff1 * x + coeff2 * y, obj1, obj2) |
|
|
|
|
| @misc.register_state_class |
| class WeightedMovingAverage(Generic[TArrayTree], misc.State): |
| """A wrapped class for an arbitrary weighted moving average.""" |
|
|
| weight: Numeric |
| value: TArrayTree | None |
|
|
| @property |
| def ndim(self) -> int: |
| assert self.value is not None |
| return self.value.ndim |
|
|
| @property |
| def shape(self) -> Shape: |
| assert self.value is not None |
| return self.value.shape |
|
|
| @property |
| def dtype(self) -> DType: |
| assert self.value is not None |
| return self.value.dtype |
|
|
| def update( |
| self, |
| value: TArrayTree, |
| old_weight_multiplier: Numeric, |
| new_weight: Numeric, |
| add_function: AddFunction = default_add_function, |
| ): |
| """Updates the underlying array and weight accordingly.""" |
|
|
| assert self.value is not None |
|
|
| |
| |
| |
| |
| self.weight = old_weight_multiplier * self.weight + jax.nn.relu(new_weight) |
| eta_for_old = jax.nn.relu(new_weight) / self.weight |
| eta_for_new = jnp.abs(new_weight) / self.weight |
|
|
| self.value = add_function(self.value, value, 1.0 - eta_for_old, eta_for_new) |
|
|
| def sync(self, pmap_axis_name: str | None): |
| """Syncs the underlying array across devices.""" |
|
|
| if self.value is None: |
| raise ValueError("`_value` has not been set yet.") |
|
|
| self.value = parallel.pmean_if_pmap(self.value, pmap_axis_name) |
|
|
| def clear(self, value_to_none: bool = False): |
| """Resets the weighted average.""" |
|
|
| self.weight = jnp.zeros_like(self.weight) |
| self.value = None if value_to_none else jnp.zeros_like(self.value) |
|
|
| def value_and_clear(self) -> TArrayTree: |
| """Retrieves the value of the weighted average and clears it.""" |
|
|
| value = self.value |
| self.clear() |
|
|
| assert value is not None |
| return value |
|
|
| @classmethod |
| def zeros_array( |
| cls, |
| shape: Shape, |
| dtype: DType | None = None, |
| ) -> "WeightedMovingAverage[Array]": |
| """Initializes a `WeightedMovingAverage` with a single array of zeros.""" |
|
|
| return cls( |
| weight=jnp.zeros([], dtype=dtype), |
| value=jnp.zeros(shape, dtype=dtype), |
| ) |
|
|
| @classmethod |
| def zeros_like(cls, value: TArrayTree) -> "WeightedMovingAverage[TArrayTree]": |
| """Initializes a `WeightedMovingAverage` with zeros structure like `value`.""" |
|
|
| return cls( |
| weight=jnp.array( |
| 0.0, dtype=types.get_float_dtype_and_check_consistency(value) |
| ), |
| value=jax.tree_util.tree_map(jnp.zeros_like, value), |
| ) |
|
|
|
|
| class MultiChunkAccumulator(Generic[TArrayTree]): |
| """Statistics accumulation, abstracted over multiple chunks.""" |
|
|
| def __init__( |
| self, |
| init_obj_value: TArrayTree | None, |
| weight: Numeric, |
| multi_device: bool, |
| ): |
| """Initializes an accumulator instance with the provided object and counter. |
| |
| Args: |
| init_obj_value: The initial value of the accumulator. |
| weight: The initial weight, which specifies how many samples are assumed |
| to have been already counted in the initial value of the accumulator. |
| multi_device: Whether the objects that are accumulated are outputs of a |
| multi-device computation (e.g. `jax.pmap`). |
| """ |
| self._accumulator = init_obj_value |
| self._weight = weight |
| self._multi_device = multi_device |
|
|
| @property |
| def accumulator(self) -> TArrayTree | None: |
| """The current value of the underlying not-normalized accumulator.""" |
| return self._accumulator |
|
|
| @property |
| def weight(self) -> Numeric | None: |
| """The current normalization weight of the underlying accumulator.""" |
| return self._weight |
|
|
| @property |
| def multi_device(self) -> bool: |
| """Whether the accumulator is the output of a multi-device computation.""" |
| return self._multi_device |
|
|
| @property |
| def value(self) -> TArrayTree | None: |
| """The current normalized value of the accumulator.""" |
|
|
| if types.tree_is_empty(self.accumulator): |
| return self.accumulator |
|
|
| if self._multi_device: |
| return parallel.pmap_sync_and_divide_value(self.accumulator, self.weight) |
| else: |
| return parallel.jit_sync_and_divide_value(self.accumulator, self.weight) |
|
|
| def clear(self) -> None: |
| """Sets the underlying accumulator and weight to `None`.""" |
| self._accumulator = None |
| self._weight = None |
|
|
| def value_and_clear(self) -> TArrayTree | None: |
| """Retrieves the normalized value of the accumulator and clears it.""" |
|
|
| value = self.value |
| self.clear() |
|
|
| return value |
|
|
| def add(self, value_obj: TArrayTree, weight: Numeric = 1): |
| """Adds an element to the moving average and the max. |
| |
| The exact update equation for the statistics are: |
| raw_value_t = raw_value_{t-1} + value_obj * weight |
| weight_t = weight_{t-1} + weight |
| |
| Args: |
| value_obj: The value of the object, which scaled by `weight` will be added |
| to the accumulator. |
| weight: The relative weight of the `value_obj`. |
| """ |
|
|
| value_obj = jax.tree_util.tree_map(lambda x: x * weight, value_obj) |
|
|
| if self._accumulator is None: |
|
|
| self._accumulator = value_obj |
|
|
| if isinstance(weight, types.SCALAR_TYPES): |
| self._weight = jnp.full_like(self._weight, weight) |
|
|
| elif not isinstance(weight, jax.Array): |
| raise ValueError("`weight` should be an instance of float, int or " |
| "jax.Array.") |
|
|
| elif self._weight.shape != weight.shape: |
| raise ValueError("If `weight` is an `jnp.ndarray` then should have the " |
| "same shape as the weight of the accumulator.") |
| else: |
| self._weight = weight |
|
|
| return |
|
|
| if not types.tree_is_empty(self._accumulator): |
|
|
| if types.tree_is_empty(value_obj): |
| raise ValueError("The provided `value_obj` has an empty PyTree " |
| "structure, but the accumulator has been initialized " |
| "with a non-empty PyTree object.") |
|
|
| self._accumulator = jax.tree_util.tree_map( |
| jnp.add, self._accumulator, value_obj) |
|
|
| elif not types.tree_is_empty(value_obj): |
|
|
| raise ValueError("The provided `value_obj` has a non-empty PyTree " |
| "structure, but the accumulator has been initialized " |
| "with an empty PyTree object.") |
|
|
| self._weight = self._weight + weight |
|
|
| @classmethod |
| def zeros_like( |
| cls, |
| obj: TArrayTree, |
| multi_device: bool |
| ) -> "MultiChunkAccumulator[TArrayTree]": |
| """Creates a zero initialized accumulator as `obj`.""" |
|
|
| if multi_device: |
| value = (parallel.pmap_zeros_like(obj) |
| if not types.tree_is_empty(obj) else obj) |
| weight = parallel.replicate_all_local_devices( |
| jnp.zeros([], dtype=jnp.int32)) |
| else: |
| value = (parallel.jit_zeros_like(obj) |
| if not types.tree_is_empty(obj) else obj) |
| weight = jnp.zeros([], dtype=jnp.int32) |
|
|
| return cls(value, weight, multi_device) |
|
|
| @classmethod |
| def empty(cls, multi_device: bool) -> "MultiChunkAccumulator[Any]": |
| """Creates an empty accumulator.""" |
|
|
| weight = jnp.zeros([], dtype=jnp.int32) |
|
|
| if multi_device: |
| weight = parallel.replicate_all_local_devices(weight) |
|
|
| return cls(None, weight, multi_device) |
|
|
| def __repr__(self): |
| return (f"{self.__class__.__name__}({self._accumulator!r}, " |
| f"{self._weight!r}, {self._multi_device})") |
|
|
| def copy(self): |
| """Returns a copy of the PyTree structure (but not the JAX arrays).""" |
|
|
| (flattened, structure) = jax.tree_util.tree_flatten(self) |
|
|
| return jax.tree_util.tree_unflatten(structure, flattened) |
|
|
|
|
| jax.tree_util.register_pytree_node( |
| MultiChunkAccumulator, |
| lambda x: ((x.accumulator, x.weight), (x.multi_device,)), |
| lambda fixed, arrays: MultiChunkAccumulator(*arrays, *fixed) |
| ) |
|
|