| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| """The kfac_jax optimizer (supporting K-FAC and other methods).""" |
|
|
| import functools |
| from typing import Any, Callable, Generic, Iterator, Sequence |
| from absl import logging |
|
|
| import jax |
| from jax import lax |
| import jax.numpy as jnp |
| from kfac_jax._src import curvature_estimator |
| from kfac_jax._src import utils |
| from typing_extensions import Self |
|
|
|
|
| |
| Array = utils.Array |
| PRNGKey = utils.PRNGKey |
| Numeric = utils.Numeric |
| Params = utils.Params |
| Batch = utils.Batch |
| FuncState = Any |
| FuncAux = utils.FuncAux |
| Scalar = utils.Scalar |
| ScheduleType = utils.ScheduleType |
|
|
| FuncArgsVariants = ( |
| tuple[Params, Batch] | |
| tuple[Params, FuncState, Batch] | |
| tuple[Params, PRNGKey, Batch] | |
| tuple[Params, FuncState, PRNGKey, Batch] |
| ) |
| FuncOutputs = ( |
| Array | |
| tuple[Array, FuncState] | |
| tuple[Array, FuncAux] | |
| tuple[Array, tuple[FuncState, FuncAux]] |
| ) |
| ValueFunc = Callable[..., FuncOutputs] |
| ValueAndGradFunc = Callable[..., tuple[FuncOutputs, Params]] |
| SharedForwardFunc = Callable[..., tuple[Array, Array]] |
| BlockDiagonalCurvature = curvature_estimator.BlockDiagonalCurvature |
|
|
| ReturnEither = ( |
| tuple[Params, "Optimizer.State", FuncState, dict[str, Numeric]] | |
| tuple[Params, "Optimizer.State", dict[str, Numeric]] |
| ) |
|
|
| QuadModelParams = tuple[Array, Array, Array, Array] |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
|
|
| |
| |
| |
| |
| |
|
|
| HAIKU_BIASES = "b,bias" |
| HAIKU_BIASES_AND_NORMS = "b,bias,scale,offset" |
|
|
|
|
| class Optimizer(utils.WithStagedMethods): |
| """The kfac_jax optimizer (supporting K-FAC and other methods).""" |
|
|
| @utils.register_state_class |
| class State(Generic[Params], utils.State): |
| r"""Persistent state of the optimizer. |
| |
| Attributes: |
| velocities: The update to the parameters from the previous step - |
| :math:`\theta_t - \theta_{t-1}`. |
| estimator_state: The persistent state for the curvature estimator. |
| damping: When using damping adaptation, this will contain the current |
| value. |
| data_seen: The number of training cases that the optimizer has processed. |
| step_counter: An integer giving the current step number :math:`t`. |
| """ |
| velocities: Params |
| estimator_state: BlockDiagonalCurvature.State |
| damping: Array |
| data_seen: Numeric |
| step_counter: Numeric |
|
|
| @classmethod |
| def from_dict(cls, dict_representation: dict[str, Any]) -> Self: |
| dict_representation["estimator_state"] = ( |
| BlockDiagonalCurvature.State.from_dict( |
| dict_representation["estimator_state"] |
| ) |
| ) |
| return cls(**dict_representation) |
|
|
| def __init__( |
| self, |
| value_and_grad_func: ValueAndGradFunc, |
| l2_reg: Numeric, |
| regularized_parameters_path_exclusions: str = "", |
| value_func_has_aux: bool = False, |
| value_func_has_state: bool = False, |
| value_func_has_rng: bool = False, |
| value_func_for_estimator: ValueFunc | None = None, |
| use_adaptive_learning_rate: bool = False, |
| learning_rate_schedule: ScheduleType | None = None, |
| use_adaptive_momentum: bool = False, |
| momentum_schedule: ScheduleType | None = None, |
| use_adaptive_damping: bool = False, |
| damping_schedule: ScheduleType | None = None, |
| initial_damping: Numeric | None = None, |
| use_initial_damping_calibration: bool = False, |
| min_damping: Numeric = 1e-8, |
| max_damping: Numeric = jnp.inf, |
| include_damping_in_quad_change: bool = False, |
| damping_adaptation_interval: int = 5, |
| damping_adaptation_decay: Numeric = 0.9, |
| damping_lower_threshold: Numeric = 0.25, |
| damping_upper_threshold: Numeric = 0.75, |
| always_use_exact_qmodel_for_damping_adjustment: bool = False, |
| precon_damping_mult: Numeric = 1.0, |
| precon_damping_schedule: ScheduleType | None = None, |
| use_step_rejection: bool = False, |
| reject_damping_increase_factor: float = 1.0, |
| norm_constraint: Numeric | None = None, |
| num_burnin_steps: int = 10, |
| estimation_mode: str | None = None, |
| custom_estimator_ctor: ( |
| Callable[..., BlockDiagonalCurvature] | None) = None, |
| curvature_ema: Numeric = 0.95, |
| curvature_update_period: int = 1, |
| inverse_update_period: int = 5, |
| use_exact_inverses: bool = False, |
| batch_process_func: Callable[[Batch], Batch] | None = None, |
| register_only_generic: bool = False, |
| patterns_to_skip: Sequence[str] = (), |
| use_automatic_registration: bool = True, |
| auto_register_kwargs: dict[str, Any] | None = None, |
| layer_tag_to_block_ctor: ( |
| dict[str, curvature_estimator.CurvatureBlockCtor] | None) = None, |
| multi_device: bool = False, |
| debug: bool = False, |
| invalid_metric_value: Numeric = jnp.nan, |
| batch_size_extractor: Callable[ |
| [Batch], Numeric |
| ] = utils.default_batch_size_extractor, |
| pmap_axis_name: str = "batch_axis", |
| forbid_setting_attributes_after_finalize: bool = True, |
| modifiable_attribute_exceptions: Sequence[str] = (), |
| include_norms_in_stats: bool = False, |
| include_per_param_norms_in_stats: bool = False, |
| include_registered_loss_in_stats: bool = False, |
| distributed_precon_apply: bool = True, |
| distributed_inverses: bool = True, |
| num_estimator_samples: int = 1, |
| should_vmap_estimator_samples: bool = False, |
| norm_to_scale_identity_weight_per_block: str | None = None, |
| step_stats_hook: Callable[..., dict[str, Array]] | None = None, |
| precon_power: Scalar = -1.0, |
| exact_quad_model_matrix_type: str | None = None, |
| value_func_for_shared_forward: SharedForwardFunc | None = None, |
| share_curvature_and_grad_forward: bool = False, |
| ): |
| """Initializes the kfac_jax optimizer with the provided settings. |
| |
| NOTE: Please read the docstring for this constructor carefully. Especially |
| the description of ``value_and_grad_func``. |
| |
| A note on the "damping" parameter: |
| |
| One of the main complications of using second-order optimizers like K-FAC is |
| the "damping" parameter. This parameter is multiplied by the identity matrix |
| and (approximately) added to the curvature matrix (i.e. the Fisher or GGN) |
| before it is inverted and multiplied by the gradient when computing the |
| update (before any learning rate scaling). The damping should follow the |
| scale of the objective, so that if you multiply your loss by some factor you |
| should do the same for the damping. Roughly speaking, larger damping values |
| constrain the update vector to a smaller region around zero, which is needed |
| in general since the second-order approximations that underlie second-order |
| methods can break down for large updates. (In gradient descent the learning |
| rate plays an analogous role.) The relationship between the damping |
| parameter and the radius of this region is complicated and depends on the |
| scale of the objective amongst other things. |
| |
| The optimizer provides a system for adjusting the damping automatically via |
| the ``use_adaptive_damping`` argument, although this system is not reliable, |
| especially for highly stochastic objectives. Using a fixed value or a |
| manually tuned schedule can work as good or better for some problems, while |
| it can be a very poor choice for others (like deep autoencoders). |
| Empirically we have found that using a fixed value works well enough for |
| common architectures like convnets and transformers. |
| |
| Args: |
| value_and_grad_func: Python callable. This function should return the |
| value of the loss to be optimized and its gradients, and optionally the |
| model state and auxiliary information in the form of a a dict mapping |
| strings to scalar arrays (usually statistics to log). Note that it |
| should *not* be jitted/pmapped or otherwise compiled by JAX, as this can |
| lead to errors. (Compilation is done internally by the optimizer.) The |
| interface of this function should be: ``out_args, loss_grads = |
| value_and_grad_func(*in_args)``. Here, ``in_args`` is ``(params, |
| func_state, rng, batch)``, with ``rng`` omitted if |
| ``value_func_has_rng`` is ``False``, and with ``func_state`` omitted if |
| ``value_func_has_state`` is ``False``. Meanwhile, ``out_args`` is |
| ``(loss, (func_state, aux))`` if ``value_func_has_state`` and |
| ``value_func_has_aux`` are both ``True``, ``(loss, func_state)`` if |
| ``value_func_has_state`` is ``True`` and ``value_func_has_aux`` is |
| ``False``, ``(loss, aux)`` if ``value_func_has_state`` is ``False`` and |
| ``value_func_has_aux`` is ``True``, and finally ``loss`` if |
| ``value_func_has_state`` and ``value_func_has_aux`` are both ``False``. |
| This should be consistent with how JAX's ``value_and_grad`` API function |
| is typically used. Note that the value (and its gradient) should be |
| normalized by the batch size, as is standard convention. Additional |
| normalization, such as by the sequence length, is up to the user, but |
| must by properly reported in the loss registration (by setting the |
| ``weight`` arguments in the loss registration functions.) |
| l2_reg: Scalar. Set this value to tell the optimizer what L2 |
| regularization coefficient you are using (if any). Note the coefficient |
| appears in the regularizer as ``coeff / 2 * sum(param**2)``. This adds |
| an additional diagonal term to the curvature and hence will affect the |
| quadratic model when using adaptive damping. Note that the user is still |
| responsible for adding regularization to the loss. |
| regularized_parameters_path_exclusions: str. A comma-separated list |
| specifying the names of parameters that should not be regularized. |
| A number of convenience examples are given in this module, e.g. |
| HAIKU_BIASES_AND_NORMS, which is ``"b,bias,scale,offset"``. |
| (Default: ``""``) |
| value_func_has_aux: Boolean. Specifies whether the provided callable |
| ``value_and_grad_func`` returns auxiliary data. (Default: ``False``) |
| value_func_has_state: Boolean. Specifies whether the provided callable |
| ``value_and_grad_func`` has a persistent state that is passed in and |
| out. (Default: ``False``) |
| value_func_has_rng: Boolean. Specifies whether the provided callable |
| ``value_and_grad_func`` additionally takes as input an rng key. |
| (Default: ``False``) |
| value_func_for_estimator: ValueFunc. If specified, this function will be |
| used by the preconditioner estimator instead of ``value_and_grad_func``. |
| This is useful for cases where the value function used for training is |
| expensive to add to the preconditioner, e.g. because it has costly |
| regularizers. (Default: ``None``) |
| value_func_for_shared_forward: Tagged function returning |
| ``(loss, gradient_surrogate)``. The surrogate's parameter gradient |
| must equal the training gradient. Required when |
| ``share_curvature_and_grad_forward=True``. |
| use_adaptive_learning_rate: Boolean. Specifies whether to use the special |
| rule from the original K-FAC paper for picking the learning rate at each |
| step. Note that this won't work well for stochastic objectives. If this |
| is ``False``, the user must use the ``learning_rate`` argument of the |
| step function, or the constructor argument ``learning_rate_schedule``. |
| (Default: ``False``) |
| learning_rate_schedule: Callable. A schedule for the learning rate. This |
| should take as input the current step number, and optionally the amount |
| of data seen so far as a keyword argument ``data_seen``, and return a |
| single array that represents the learning rate. (Default: ``None``) |
| use_adaptive_momentum: Boolean. Specifies whether to use the special rule |
| from the original K-FAC paper for picking the momentum "decay" parameter |
| at each step. Note that this won't work well for stochastic objectives. |
| If this is ``False``, the user must use the ``momentum`` argument of the |
| step function, or the constructor argument ``momentum_schedule``. |
| (Default: ``False``) |
| momentum_schedule: Callable. A schedule for the momentum parameter. This |
| should take as input the current step number, and optionally the amount |
| of data seen so far as a keyword argument ``data_seen``, and return a |
| single array that represents the momentum. (Default: ``None``) |
| use_adaptive_damping: Boolean. Specifies whether the optimizer will use |
| the Levenberg-Marquardt method to automatically adjust the damping every |
| ``damping_adaptation_interval`` iterations. If this is set to ``False`` |
| the user must provide a value to the damping argument of the step |
| function at each iteration, or use the ``damping_schedule`` constructor |
| argument. Note that the effectiveness of this technique seems to vary |
| between problems. (Default: ``False``) |
| damping_schedule: Callable. A schedule for the damping. This should take |
| as input the current step number, and optionally the amount of data seen |
| so far as a keyword argument ``data_seen``, and return a single array |
| that represents the learning rate. (Default: ``None``) |
| initial_damping: Scalar or None. This specifies the initial value of the |
| damping that the optimizer will use when using automatic damping |
| adaptation. (Default: ``None``) |
| use_initial_damping_calibration: Boolean. If ``True``, the initial damping |
| value, used to initialize the adaptive damping method, will be first |
| calibrated (after any burnin steps to estimate the preconditioner) so |
| that its value wouldn't be changed after the first step of optimization. |
| This calibration is done by essentially running the step function |
| multiple times without actually updating the parameters or sampling a |
| new mini-batch. ``num_burnin_steps`` must be greater than 0 to use this |
| option. (Default: ``False``) |
| min_damping: Scalar. Minimum value the damping parameter can take when |
| using automatic damping adaptation. Note that the default value of 1e-8 |
| is quite arbitrary, and you may have to adjust this up or down for your |
| particular problem. If you are using a non-zero value of l2_reg you |
| *may* be able to set this to zero. (Default: ``1e-8``) |
| max_damping: Scalar. Maximum value the damping parameter can take when |
| using automatic damping adaptation. (Default: ``Infinity``) |
| include_damping_in_quad_change: Boolean. Whether to include the |
| contribution of the damping in the quadratic model for the purposes |
| computing the reduction ration ("rho") in the Levenberg-Marquardt scheme |
| used for adapting the damping. Note that the contribution from the |
| ``l2_reg`` argument is always included. (Default: ``False``) |
| damping_adaptation_interval: Int. The number of steps in between adapting |
| the damping parameter. (Default: ``5``) |
| damping_adaptation_decay: Scalar. The damping parameter will be adjusted |
| up or down by ``damping_adaptation_decay ** |
| damping_adaptation_interval``, or remain unchanged, every |
| ``damping_adaptation_interval`` number of iterations. (Default: ``0.9``) |
| damping_lower_threshold: Scalar. The damping parameter is increased if the |
| reduction ratio is below this threshold. (Default: ``0.25``) |
| damping_upper_threshold: Scalar. The damping parameter is decreased if the |
| reduction ratio is below this threshold. (Default: ``0.75``) |
| always_use_exact_qmodel_for_damping_adjustment: Boolean. When using |
| learning rate and/or momentum adaptation, the quadratic model change |
| used for damping adaption is always computed using the exact curvature |
| matrix. Otherwise, there is an option to use either the exact or |
| approximate curvature matrix to compute the quadratic model change, |
| which is what this argument controls. When True, the exact curvature |
| matrix will be used, which is more expensive, but could possibly produce |
| a better damping schedule. (Default: ``False``) |
| precon_damping_mult: Scalar. When ``precon_damping_schedule`` is unset, |
| the regular damping is used for the preconditioner damping, multiplied |
| by this value. (Default: ``1.0``) |
| precon_damping_schedule: Similar to ``damping_schedule``, but for the |
| preconditioner only. If ``None``, the preconditioner will use the |
| regular damping, multiplied by ``precon_damping_mult``. |
| (Default: ``None``) |
| use_step_rejection: Whether or not to reject the step whenever the loss |
| on the current batch goes up after the update. This option offers |
| robustness at the cost of doing more work per step (unless adaptive |
| damping with Levenberg-Marquardt is used). (Default: ``False``) |
| reject_damping_increase_factor: The damping parameter is increased by this |
| factor if the step is rejected. (Default: ``1.0``) |
| norm_constraint: Scalar. If specified, the update is scaled down so that |
| its approximate squared Fisher norm ``v^T F v`` is at most the specified |
| value. (Note that here ``F`` is the approximate curvature matrix, not |
| the exact.) May only be used when ``use_adaptive_learning_rate`` is |
| ``False``. (Default: ``None``) |
| num_burnin_steps: Int. At the start of optimization, e.g. the first step, |
| before performing the actual step the optimizer will perform this many |
| times updates to the curvature approximation without updating the actual |
| parameters. (Default: ``10``) |
| estimation_mode: String. The type of estimator to use for the curvature |
| matrix. See the documentation for :class:`~BlockDiagonalCurvature` for a |
| detailed description of the possible options. If ``None`` will use |
| default estimation_mode mode of the used CurvatureEstimator subclass, |
| which is typically "ggn_curvature_prop". (Default: ``None``) |
| custom_estimator_ctor: Optional constructor for subclass of |
| :class:`~BlockDiagonalCurvature`. If specified, the optimizer will use |
| this conastructor instead of the default |
| :class:`~BlockDiagonalCurvature`. (Default: ``None``) |
| curvature_ema: The decay factor used when calculating the covariance |
| estimate moving averages. (Default: ``0.95``) |
| curvature_update_period: Int. The number of steps in between updating the |
| the curvature estimates. (Default: ``1``) |
| inverse_update_period: Int. The number of steps in between updating the |
| the computation of the inverse curvature approximation. (Default: ``5``) |
| use_exact_inverses: Bool. If ``True``, preconditioner inverses are |
| computed "exactly" without the pi-adjusted factored damping approach. |
| Note that this involves the use of eigendecompositions, which can |
| sometimes be much more expensive. (Default: ``False``) |
| batch_process_func: Callable. A function which to be called on each batch |
| before feeding to the KFAC on device. This could be useful for specific |
| device input optimizations. (Default: ``None``) |
| register_only_generic: Boolean. Whether when running the auto-tagger to |
| register only generic parameters, or allow it to use the graph matcher |
| to automatically pick up any kind of layer tags. (Default: ``False``) |
| patterns_to_skip: tuple. A list of any patterns that should be skipped by |
| the graph matcher when auto-tagging. (Default: ``()``) |
| use_automatic_registration: Bool. If ``True``, the optimizer will try to |
| automatically register the layers of your network. (Default: ``True``) |
| auto_register_kwargs: Any additional kwargs to be passed down to |
| :func:`~auto_register_tags`, which is called by the curvature estimator. |
| (Default: ``None``) |
| layer_tag_to_block_ctor: dictionary. A mapping from layer tags to block |
| classes which to override the default choices of block approximation for |
| that specific tag. See the documentation for |
| :class:`~CurvatureEstimator` for a more detailed description. (Default: |
| ``None``) |
| multi_device: Boolean. Whether to use pmap and run the optimizer on |
| multiple devices. (Default: ``False``) |
| debug: Boolean. If neither the step or init functions should be jitted. |
| Note that this also overrides ``multi_device`` and prevents using pmap, |
| instead using a "simulated pmap" that loops over the device index and |
| does everything on the default device. (Default: ``False``) |
| invalid_metric_value: Numeric. Certain metrics returned from the step |
| function are not always computed at each iteration, or may otherwise |
| be invalid. In such cases we need to return a value anyway. jnp.nan is |
| a natural choice, but can sometimes cause problems (e.g. false positives |
| JAX's automatic NaN checker). This argument allows the user to specify a |
| different value to return in such cases. (Default: ``jnp.nan``) |
| batch_size_extractor: A function that takes as input the function |
| arguments and returns the batch size for a single device. (Default: |
| ``kfac.utils.default_batch_size_extractor``) |
| pmap_axis_name: String. The name of the pmap axis to use when |
| ``multi_device`` is set to True. (Default: ``batch_axis``) |
| forbid_setting_attributes_after_finalize: Boolean. By default, after the |
| object is finalized, you can not set any of its properties. This is done |
| in order to protect the user from making changes to the object |
| attributes that would not be picked up by various internal methods after |
| they have been compiled. However, if you are extending this class, and |
| clearly understand the risks of modifying attributes, setting this to |
| ``False`` will remove the restriction. (Default: ``True``) |
| modifiable_attribute_exceptions: Sequence of strings. Gives a list of |
| names for attributes that can be modified after finalization even when |
| ``forbid_setting_attributes_after_finalize`` is ``True``. (Default: |
| ``()``) |
| include_norms_in_stats: Boolean. It True, the vector norms of the |
| gradient, preconditioned gradient, and parameter update are included in |
| the statistics returned by the step function. (Default: ``False``) |
| include_per_param_norms_in_stats: Boolean. It True, the per-parameter |
| vector norms of the gradient, preconditioned gradient, and parameter |
| update are included in the statistics returned by the step function. |
| (Default: ``False``) |
| include_registered_loss_in_stats: Boolean. If True, we include the loss, |
| as computed from the registered losses, in the stats. Also included is |
| the relative difference between this as the loss computed from |
| ``value_and_grad_func``. This is useful for debugging registration |
| errors. Note this for this option to work it's required that the targets |
| are passed for each loss function registration. (Default: ``False``) |
| distributed_precon_apply: Boolean. Whether to distribute the application |
| of the preconditioner across the different devices in a layer-wise |
| fashion. If False, each device will (redundantly) perform the required |
| operations for all the layers. (Default: True) |
| distributed_inverses: Boolean. Whether to distribute the inverse |
| computations (required to compute the preconditioner) across the |
| different devices in a layer-wise fashion. If False, each device will |
| (redundantly) perform the required computations for all the layers. |
| (Default: True) |
| num_estimator_samples: Number of samples (per case) to use when computing |
| stochastic curvature matrix estimates. This option is only used when |
| ``estimation_mode == 'fisher_gradients'`` or ``estimation_mode == |
| '[fisher,ggn]_curvature_prop'``. (Default: 1) |
| should_vmap_estimator_samples: Whether to use ``jax.vmap`` to compute |
| samples when ``num_estimator_samples > 1``. (Default: False) |
| share_curvature_and_grad_forward: Reuse the exact-Fisher tagged model |
| primal evaluation for the ordinary loss and training gradient on |
| curvature-update steps. (Default: ``False``) |
| norm_to_scale_identity_weight_per_block: The name of a norm to use to |
| compute extra per-block scaling for the damping. See psd_matrix_norm() |
| in utils/math.py for the definition of these. Note that this will not |
| affect the exact quadratic model that is used as part of the "adaptive" |
| learning rate, momentum, and damping methods. (Default: None) |
| step_stats_hook: Optional callable ``(estimator, grads, |
| preconditioned_gradient) -> dict`` invoked inside ``_step`` on the |
| PRE-norm-constraint preconditioned gradient; returned scalars are |
| merged into the step stats dict. Runs inside the step's jit — must |
| be trace-safe and cheap. (Default: None) |
| precon_power: The matrix power to use when computing the preconditioner. |
| K-FAC use -1 by default, but ``kfac_jax`` can simulate other optimizers |
| like RMSProp by using -0.5 (along with appropriate changes to |
| ``layer_tag_to_block_ctor`` and ``estimation_mode``). (Default: -1) |
| exact_quad_model_matrix_type: The type of matrix to use when computing the |
| exact quadratic model (used in the adaptive learning rate and momentum). |
| Can be ``'fisher'``, ``'ggn'``, or None. If None, will use the value |
| implied by ``estimation_mode``. (Default: None) |
| """ |
|
|
| super().__init__( |
| multi_device=multi_device, |
| pmap_axis_name=pmap_axis_name if multi_device else None, |
| debug=debug, |
| forbid_setting_attributes_after_finalize= |
| forbid_setting_attributes_after_finalize, |
| excluded_attribute_names=modifiable_attribute_exceptions, |
| ) |
|
|
| if use_adaptive_damping and initial_damping is None: |
| raise ValueError("When use_adaptive_damping is True you must provide a " |
| "value for initial_damping.") |
| if use_adaptive_learning_rate and learning_rate_schedule is not None: |
| raise ValueError("If you are using adaptive learning rate then " |
| "`learning_rate_schedule` should be None.") |
| if use_adaptive_momentum and momentum_schedule is not None: |
| raise ValueError("If you are using adaptive momentum then " |
| "`momentum_schedule` should be None.") |
| if use_adaptive_damping and damping_schedule is not None: |
| raise ValueError("If you are using adaptive damping then " |
| "`damping_schedule` should be None.") |
|
|
| if num_burnin_steps <= 0 and use_initial_damping_calibration: |
| raise ValueError("num_burnin_steps must be > 0 if " |
| "use_initial_damping_calibration is True.") |
|
|
| self._value_and_grad_func = value_and_grad_func |
| self._value_func_has_aux = value_func_has_aux |
| self._value_func_has_state = value_func_has_state |
| self._value_func_has_rng = value_func_has_rng |
| if share_curvature_and_grad_forward: |
| incompatible = [] |
| if estimation_mode != "fisher_exact": |
| incompatible.append("estimation_mode must be 'fisher_exact'") |
| if value_func_for_estimator is not None: |
| incompatible.append("value_func_for_estimator must be None") |
| if value_func_for_shared_forward is None: |
| incompatible.append("value_func_for_shared_forward must be provided") |
| if custom_estimator_ctor is not None: |
| incompatible.append("custom_estimator_ctor must be None") |
| if value_func_has_aux: |
| incompatible.append("value_func_has_aux must be False") |
| if value_func_has_state: |
| incompatible.append("value_func_has_state must be False") |
| if value_func_has_rng: |
| incompatible.append("value_func_has_rng must be False") |
| if include_registered_loss_in_stats: |
| incompatible.append( |
| "include_registered_loss_in_stats must be False" |
| ) |
| if incompatible: |
| raise ValueError( |
| "`share_curvature_and_grad_forward=True` is incompatible with: " |
| + "; ".join(incompatible) |
| + "." |
| ) |
| self._share_curvature_and_grad_forward = ( |
| share_curvature_and_grad_forward |
| ) |
| self._value_func: ValueFunc = convert_value_and_grad_to_value_func( |
| value_and_grad_func, |
| has_aux=value_func_has_aux or value_func_has_state, |
| ) |
|
|
| self._l2_reg = l2_reg |
| self._regularized_parameters_path_exclusions = ( |
| regularized_parameters_path_exclusions.split(",")) |
|
|
| self._use_adaptive_learning_rate = use_adaptive_learning_rate |
| self._learning_rate_schedule = learning_rate_schedule |
| self._use_adaptive_momentum = use_adaptive_momentum |
| self._momentum_schedule = momentum_schedule |
|
|
| self._use_adaptive_damping = use_adaptive_damping |
| self._damping_schedule = damping_schedule |
| self._initial_damping = initial_damping |
| self._use_initial_damping_calibration = use_initial_damping_calibration |
| self._min_damping = min_damping |
| self._max_damping = max_damping |
| self._include_damping_in_quad_change = include_damping_in_quad_change |
| self._damping_adaptation_decay = damping_adaptation_decay |
| self._damping_adaptation_interval = damping_adaptation_interval |
| self._damping_lower_threshold = damping_lower_threshold |
| self._damping_upper_threshold = damping_upper_threshold |
| self._always_use_exact_qmodel_for_damping_adjustment = ( |
| always_use_exact_qmodel_for_damping_adjustment) |
| self._precon_damping_mult = precon_damping_mult |
| self._precon_damping_schedule = precon_damping_schedule |
|
|
| self._use_step_rejection = use_step_rejection |
| self._reject_damping_increase_factor = reject_damping_increase_factor |
|
|
| self._norm_constraint = norm_constraint |
| self._num_burnin_steps = num_burnin_steps |
| self._curvature_ema = curvature_ema |
| if curvature_update_period > inverse_update_period: |
| raise ValueError( |
| "curvature_update_period ({}) cannot be larger than" |
| " inverse_update_period ({}) as the identical matrix inversion would" |
| " be redundantly performed. Set inverse_update_period larger instead." |
| .format(curvature_update_period, inverse_update_period) |
| ) |
| self._curvature_update_period = curvature_update_period |
| self._inverse_update_period = inverse_update_period |
| self._layer_tag_to_block_cls = layer_tag_to_block_ctor |
| self._patterns_to_skip = patterns_to_skip |
| self._batch_process_func = batch_process_func or (lambda x: x) |
| self._include_norms_in_stats = include_norms_in_stats |
| self._include_per_param_norms_in_stats = include_per_param_norms_in_stats |
| self._include_registered_loss_in_stats = include_registered_loss_in_stats |
| self._batch_size_extractor = batch_size_extractor |
|
|
| self.__invalid_metric_value = invalid_metric_value |
|
|
| self._use_cached_inverses = (self._inverse_update_period != 1) |
| self._use_exact_inverses = use_exact_inverses |
|
|
| self._norm_to_scale_identity_weight_per_block = ( |
| norm_to_scale_identity_weight_per_block |
| ) |
|
|
| |
| |
| |
| |
| |
| |
| self._step_stats_hook = step_stats_hook |
|
|
| self._precon_power = precon_power |
|
|
| self._exact_quad_model_matrix_type = exact_quad_model_matrix_type |
|
|
| self._params_index = 0 |
| batch_index = int(value_func_has_state + value_func_has_rng + 1) |
|
|
| if (norm_to_scale_identity_weight_per_block is not None |
| and norm_to_scale_identity_weight_per_block != "none"): |
|
|
| assert (not use_adaptive_learning_rate and not use_adaptive_momentum |
| and not use_adaptive_damping) |
|
|
| estimator_ctor = (custom_estimator_ctor or BlockDiagonalCurvature) |
|
|
| auto_register_kwargs = auto_register_kwargs or {} |
| auto_register_kwargs.update(dict( |
| register_only_generic=register_only_generic, |
| patterns_to_skip=patterns_to_skip, |
| )) |
|
|
| if value_func_for_estimator is None: |
| |
| |
| |
| |
| |
| |
| |
| func_and_grad_for_estimator = convert_value_and_grad_to_clean_value_and_grad( |
| value_and_grad_func, |
| has_aux=value_func_has_aux or value_func_has_state, |
| ) |
|
|
| else: |
| func_and_grad_for_estimator = None |
|
|
| estimator_extra_kwargs = {} |
| if share_curvature_and_grad_forward: |
| estimator_extra_kwargs["shared_forward_value_func"] = ( |
| value_func_for_shared_forward |
| ) |
|
|
| |
| self._estimator = estimator_ctor( |
| func=value_func_for_estimator, |
| func_and_grad=func_and_grad_for_estimator, |
| default_estimation_mode=estimation_mode, |
| params_index=self._params_index, |
| batch_index=batch_index, |
| layer_tag_to_block_ctor=layer_tag_to_block_ctor, |
| distributed_multiplies=distributed_precon_apply, |
| distributed_cache_updates=distributed_inverses, |
| num_samples=num_estimator_samples, |
| should_vmap_samples=should_vmap_estimator_samples, |
| auto_register_tags=use_automatic_registration, |
| auto_register_kwargs=auto_register_kwargs, |
| **estimator_extra_kwargs, |
| ) |
| self._implicit = curvature_estimator.ImplicitExactCurvature( |
| self._value_func, |
| params_index=self._params_index, |
| batch_size_extractor=batch_size_extractor, |
| ) |
|
|
| |
| |
| if type(self) == Optimizer: |
| self.finalize() |
|
|
| @property |
| def _invalid_metric_value(self) -> Array: |
| return jnp.array(self.__invalid_metric_value, dtype=float) |
|
|
| @property |
| def _damping_decay_factor(self) -> Numeric: |
| """How fast to decay the damping, when using damping adaptation.""" |
| return self._damping_adaptation_decay ** self._damping_adaptation_interval |
|
|
| @property |
| def _exact_powers_to_cache(self) -> Numeric | Sequence[Numeric] | None: |
| if self._use_exact_inverses and self._use_cached_inverses: |
| return self._precon_power |
| else: |
| return None |
|
|
| @property |
| def _approx_powers_to_cache(self) -> Numeric | Sequence[Numeric] | None: |
| if not self._use_exact_inverses and self._use_cached_inverses: |
| return self._precon_power |
| else: |
| return None |
|
|
| @property |
| def _mat_type_for_exact_quad_model(self) -> str: |
| if self._exact_quad_model_matrix_type is None: |
| return self._estimator.default_mat_type |
| return self._exact_quad_model_matrix_type |
|
|
| def _should_update_damping(self, step_counter: int) -> bool: |
| """Whether at the current step the optimizer should update the damping.""" |
| return ((step_counter + 1) % self._damping_adaptation_interval == 0) and ( |
| self._use_adaptive_damping |
| ) |
|
|
| def _should_update_estimate_curvature(self, step_counter: int) -> bool: |
| """Whether at the current step the optimizer should update the curvature estimates.""" |
| return step_counter % self._curvature_update_period == 0 |
|
|
| def _should_update_inverse_cache( |
| self, |
| state: State, |
| inverse_update_period: Numeric | None = None, |
| ) -> Array | bool: |
| """Whether at the current step the optimizer should update the inverse curvature approximation.""" |
| period = (self._inverse_update_period if inverse_update_period is None |
| else inverse_update_period) |
| return self._use_cached_inverses and ( |
| state.step_counter % period == 0) |
|
|
| def _should_sync_estimator( |
| self, |
| state: State, |
| inverse_update_period: Numeric | None = None, |
| ) -> Array | bool: |
| """Whether at the current step the optimizer should update the inverse curvature approximation.""" |
|
|
| if self._use_cached_inverses: |
| return self._should_update_inverse_cache(state, inverse_update_period) |
|
|
| return True |
|
|
| def set_live_hparams( |
| self, |
| *, |
| curvature_ema: Numeric | None = None, |
| curvature_update_period: int | None = None, |
| inverse_update_period: int | None = None, |
| ) -> None: |
| """Updates cadence/EMA hyperparameters on a live optimizer instance. |
| |
| All three take effect from the next call to ``step`` without triggering |
| recompilation: ``curvature_update_period`` only enters Python-side |
| executable selection, while ``curvature_ema`` and ``inverse_update_period`` |
| are threaded into the compiled step function as runtime scalars. |
| |
| ``inverse_update_period`` cannot be changed on an optimizer constructed |
| with ``inverse_update_period=1``, since that construction permanently |
| disables the inverse cache in the estimator state. |
| """ |
| new_curv = (self._curvature_update_period if curvature_update_period is None |
| else int(curvature_update_period)) |
| new_inv = (self._inverse_update_period if inverse_update_period is None |
| else int(inverse_update_period)) |
| if new_curv < 1 or new_inv < 1: |
| raise ValueError("Update periods must be positive integers.") |
| if new_inv != self._inverse_update_period and not self._use_cached_inverses: |
| raise ValueError( |
| "Cannot change inverse_update_period on an optimizer constructed " |
| "with inverse_update_period=1 (inverse cache disabled).") |
| if new_curv > new_inv: |
| raise ValueError( |
| "curvature_update_period ({}) cannot be larger than" |
| " inverse_update_period ({}).".format(new_curv, new_inv)) |
| if curvature_ema is not None and not 0.0 <= float(curvature_ema) <= 1.0: |
| raise ValueError("curvature_ema must be in [0, 1].") |
| self.unlock_attributes() |
| try: |
| self._curvature_update_period = new_curv |
| self._inverse_update_period = new_inv |
| if curvature_ema is not None: |
| self._curvature_ema = float(curvature_ema) |
| finally: |
| self.lock_attributes() |
|
|
| def _live_step_scalars(self) -> tuple[Array, Array]: |
| """Current curvature_ema / inverse_update_period as traced step args. |
| |
| Plain rank-0 arrays, matching how callers pass learning_rate / |
| momentum / damping into ``step``: the staging layer broadcasts them, |
| so no per-device replication is required here (and |
| ``device_put_replicated`` no longer exists on modern JAX anyway). |
| """ |
| ema = jnp.asarray(self._curvature_ema, dtype=jnp.float32) |
| period = jnp.asarray(self._inverse_update_period, dtype=jnp.int32) |
| return ema, period |
|
|
| @functools.partial(utils.staged, static_argnums=1) |
| def _rng_split(self, rng: PRNGKey, num: int) -> tuple[Array, ...]: |
| """Splits the ``rng`` key.""" |
| return tuple(jax.random.split(rng, num)) |
|
|
| @utils.auto_scope_method |
| def _compute_loss_value(self, func_args: FuncArgsVariants) -> Array: |
| """Computes the value of the loss function being optimized.""" |
| return self._value_func(*func_args) |
|
|
| def _verify_args_and_get_step_counter( |
| self, |
| step_counter: Array, |
| learning_rate: Array | None = None, |
| momentum: Array | None = None, |
| damping: Array | None = None, |
| global_step_int: int | None = None, |
| ) -> int: |
| """Verifies that the arguments passed to the step function are correct.""" |
|
|
| |
| if self._use_adaptive_learning_rate and learning_rate is not None: |
| raise ValueError("When use_adaptive_learning_rate is set to True you " |
| "should not pass a value to the step function.") |
|
|
| elif not self._use_adaptive_learning_rate and ( |
| self._learning_rate_schedule is None and learning_rate is None): |
| raise ValueError("When `use_adaptive_learning_rate` is set to False and " |
| "`learning_rate_schedule` is None you must provide a " |
| "value to the step function.") |
|
|
| elif self._learning_rate_schedule is not None and learning_rate is not None: |
| raise ValueError("When you have passed a `learning_rate_schedule` you " |
| "should not pass a value to the step function.") |
|
|
| if self._use_adaptive_momentum and momentum is not None: |
| raise ValueError("When `use_adaptive_momentum` is set to True you " |
| "should not pass a value to the step function.") |
|
|
| elif not self._use_adaptive_momentum and ( |
| self._momentum_schedule is None and momentum is None): |
| raise ValueError("When `use_adaptive_momentum` is set to False and " |
| "`momentum_schedule` is None you must provide a value to" |
| " the step function.") |
|
|
| elif self._momentum_schedule is not None and momentum is not None: |
| raise ValueError("When you have passed a `momentum_schedule` you should " |
| "not pass a value to the step function.") |
|
|
| if self._use_adaptive_damping and damping is not None: |
| raise ValueError("When `use_adaptive_damping` is set to True you " |
| "should not pass a value to the step function.") |
|
|
| elif not self._use_adaptive_damping and ( |
| self._damping_schedule is None and damping is None): |
| raise ValueError("When `use_adaptive_damping` is set to False and " |
| "`damping_schedule` is None you must provide a value to " |
| "the step function.") |
|
|
| elif self._damping_schedule is not None and damping is not None: |
| raise ValueError("When you have passed a `damping_schedule` you should " |
| "not pass a value to the step function.") |
|
|
| if global_step_int is None: |
| return int(self.get_first(step_counter)) |
|
|
| return global_step_int |
|
|
| @utils.staged |
| def _setup_state_and_schedules( |
| self, |
| learning_rate: Array | None, |
| momentum: Array | None, |
| damping: Array | None, |
| step_counter: Array, |
| data_seen: Array, |
| ) -> tuple[Numeric | None, Numeric | None, Numeric, Numeric]: |
| """Helper function for setting up learning rate, momentum and damping.""" |
|
|
| |
| if self._learning_rate_schedule is not None: |
|
|
| assert learning_rate is None |
| learning_rate = utils.call_func_with_conditional_kwargs( |
| self._learning_rate_schedule, step_counter, data_seen=data_seen) |
|
|
| if self._momentum_schedule is not None: |
| assert momentum is None |
| momentum = utils.call_func_with_conditional_kwargs( |
| self._momentum_schedule, step_counter, data_seen=data_seen) |
|
|
| if self._damping_schedule is not None: |
| assert damping is None |
| damping = utils.call_func_with_conditional_kwargs( |
| self._damping_schedule, step_counter, data_seen=data_seen) |
|
|
| else: |
| assert damping is not None |
|
|
| if self._precon_damping_schedule is not None: |
| precon_damping = utils.call_func_with_conditional_kwargs( |
| self._precon_damping_schedule, step_counter, data_seen=data_seen) |
| else: |
| precon_damping = damping * self._precon_damping_mult |
|
|
| return learning_rate, momentum, damping, precon_damping |
|
|
| def _setup_func_args_and_rng( |
| self, |
| params: Params, |
| rng: PRNGKey, |
| batch: Batch, |
| func_state: FuncState | None, |
| ) -> tuple[FuncArgsVariants, Array]: |
| """Helper function for setting up the model function arguments correctly.""" |
|
|
| |
| batch = self._batch_process_func(batch) |
|
|
| |
| if self._value_func_has_rng: |
| rng, func_rng = jax.random.split(rng) |
| else: |
| func_rng = None |
|
|
| |
| func_args = make_func_args( |
| params=params, |
| func_state=func_state, |
| rng=func_rng, |
| batch=batch, |
| has_state=self._value_func_has_state, |
| has_rng=self._value_func_has_rng, |
| ) |
|
|
| return func_args, rng |
|
|
| def _update_estimator_curvature( |
| self, |
| estimator_state: BlockDiagonalCurvature.State, |
| func_args: FuncArgsVariants, |
| rng: PRNGKey, |
| ema_old: Numeric, |
| ema_new: Numeric, |
| precon_damping: Numeric, |
| sync: Array | bool = True |
| ) -> BlockDiagonalCurvature.State: |
| """Updates the curvature estimator state.""" |
|
|
| state = self._estimator.update_curvature_matrix_estimate( |
| state=estimator_state, |
| ema_old=ema_old, |
| ema_new=ema_new, |
| identity_weight=self._l2_reg + precon_damping, |
| |
| batch_size=self._batch_size_extractor(func_args[-1]), |
| rng=rng, |
| func_args=func_args, |
| pmap_axis_name=self.pmap_axis_name, |
| ) |
| return jax.lax.cond( |
| sync, |
| functools.partial(self._estimator.sync, |
| pmap_axis_name=self.pmap_axis_name), |
| lambda state_: state_, |
| state, |
| ) |
|
|
| def _update_estimator_curvature_and_value_and_grad( |
| self, |
| estimator_state: BlockDiagonalCurvature.State, |
| func_args: FuncArgsVariants, |
| rng: PRNGKey, |
| ema_old: Numeric, |
| ema_new: Numeric, |
| precon_damping: Numeric, |
| sync: Array | bool = True, |
| ) -> tuple[BlockDiagonalCurvature.State, Array, Params]: |
| """Updates exact-Fisher curvature and returns its shared loss/gradient.""" |
|
|
| state, loss, grads = ( |
| self._estimator.update_curvature_matrix_estimate_and_value_and_grad( |
| state=estimator_state, |
| ema_old=ema_old, |
| ema_new=ema_new, |
| identity_weight=self._l2_reg + precon_damping, |
| batch_size=self._batch_size_extractor(func_args[-1]), |
| rng=rng, |
| func_args=func_args, |
| pmap_axis_name=self.pmap_axis_name, |
| ) |
| ) |
| state = jax.lax.cond( |
| sync, |
| functools.partial( |
| self._estimator.sync, |
| pmap_axis_name=self.pmap_axis_name, |
| ), |
| lambda state_: state_, |
| state, |
| ) |
| return state, loss, grads |
|
|
| @utils.auto_scope_method |
| def _compute_loss_and_grads( |
| self, |
| func_args: FuncArgsVariants, |
| state: State | None = None, |
| ) -> tuple[Array, Params, FuncState | None, FuncAux | None]: |
| """Computes the model loss value and its gradients.""" |
|
|
| del state |
|
|
| out, grads = self._value_and_grad_func(*func_args) |
|
|
| loss, func_state, aux = extract_func_outputs( |
| out, self._value_func_has_aux, self._value_func_has_state) |
|
|
| if self._include_registered_loss_in_stats: |
| aux = aux or {} |
| aux["loss_registered"] = self._compute_loss_from_registrations(func_args) |
|
|
| return loss, grads, func_state, aux |
|
|
| @functools.partial(utils.staged, donate_argnums=0) |
| def _maybe_update_inverse_cache( |
| self, |
| state: State, |
| precon_damping: Array, |
| inverse_update_period: Array, |
| ) -> State: |
| """Updates the estimator state cache if it is the right iteration.""" |
|
|
| |
| state = state.copy() |
|
|
| state.estimator_state = lax.cond( |
| self._should_update_inverse_cache(state, inverse_update_period), |
| functools.partial( |
| self._estimator.update_cache, |
| identity_weight=self._l2_reg + precon_damping, |
| exact_powers=self._exact_powers_to_cache, |
| approx_powers=self._approx_powers_to_cache, |
| eigenvalues=False, |
| pmap_axis_name=self.pmap_axis_name, |
| ), |
| lambda state_: state_, |
| state.estimator_state, |
| ) |
|
|
| return state |
|
|
| @functools.partial(utils.staged, static_argnums=3) |
| def _compute_preconditioned_gradient( |
| self, |
| state: State, |
| grads: Params, |
| precon_damping: Array, |
| can_distribute: bool = True, |
| ) -> Params: |
| """Computes the preconditioned gradient.""" |
|
|
| return self._estimator.multiply_matpower( |
| state=state.estimator_state, |
| parameter_structured_vector=grads, |
| identity_weight=self._l2_reg + precon_damping, |
| power=self._precon_power, |
| exact_power=self._use_exact_inverses, |
| use_cached=self._use_cached_inverses, |
| pmap_axis_name=self.pmap_axis_name if can_distribute else None, |
| norm_to_scale_identity_weight_per_block=self._norm_to_scale_identity_weight_per_block, |
| ) |
|
|
| @utils.staged |
| def _maybe_apply_norm_constraint( |
| self, grads: Params, preconditioned_grads: Params, coefficient: Array |
| ) -> tuple[Params, Params | None]: |
| """Scales precon grad to have curvature-weighted norm <= norm_constraint.""" |
|
|
| if self._norm_constraint is None: |
| return preconditioned_grads, None |
|
|
| assert not self._use_adaptive_learning_rate |
|
|
| sq_norm_grads = utils.inner_product(preconditioned_grads, grads) |
| sq_norm_scaled_grads = sq_norm_grads * coefficient ** 2 |
|
|
| max_coefficient = jnp.sqrt(self._norm_constraint / sq_norm_scaled_grads) |
| coefficient = jnp.minimum(max_coefficient, 1) |
|
|
| precon_grad = utils.scalar_mul(preconditioned_grads, coefficient) |
|
|
| return precon_grad, sq_norm_scaled_grads |
|
|
| def _compute_quad_change_for_damping_adapt( |
| self, |
| state: State, |
| delta: Params, |
| grads: Params, |
| damping: Array, |
| func_args: FuncArgsVariants, |
| ) -> Array: |
| """The quadratic model change, when lr and momentum are non-adaptive.""" |
|
|
| assert not (self._use_adaptive_learning_rate or self._use_adaptive_momentum) |
|
|
| if self._always_use_exact_qmodel_for_damping_adjustment: |
| quad_model = self._compute_exact_quad_model_filtered( |
| [delta], grads, func_args, state=state) |
| else: |
| quad_model = self._compute_approx_quad_model(state, [delta], grads) |
|
|
| w = jnp.ones([]) |
| return self._solve_quad_model(quad_model, damping, [w])[1] |
|
|
| def _coefficients_and_quad_change( |
| self, |
| state: State, |
| vectors: Sequence[Params], |
| grads: Params, |
| learning_rate: Numeric | None, |
| momentum: Numeric | None, |
| damping: Numeric, |
| func_args: FuncArgsVariants, |
| should_update_damping: bool, |
| ) -> tuple[tuple[Numeric, Numeric], Numeric]: |
| """The correct update coefficients and corresponding quadratic change.""" |
|
|
| |
| |
| |
| |
| neg_learning_rate = -learning_rate if learning_rate is not None else None |
| fixed_coefficients = (neg_learning_rate, momentum) |
|
|
| if self._use_adaptive_learning_rate or self._use_adaptive_momentum: |
|
|
| assert fixed_coefficients[0] is None or fixed_coefficients[1] is None |
|
|
| quad_model = self._compute_exact_quad_model_filtered( |
| vectors, grads, func_args, state=state, |
| fixed_coefficients=fixed_coefficients) |
|
|
| return self._solve_quad_model(quad_model, damping, fixed_coefficients) |
|
|
| else: |
|
|
| assert all(c is not None for c in fixed_coefficients) |
| fixed_coefficients: tuple[Numeric, Numeric] |
|
|
| if should_update_damping: |
|
|
| delta = self._weighted_sum_of_objects(vectors, fixed_coefficients) |
|
|
| quad_change = self._compute_quad_change_for_damping_adapt( |
| state, delta, grads, damping, func_args) |
|
|
| else: |
| quad_change = self._invalid_metric_value |
|
|
| return fixed_coefficients, quad_change |
|
|
| @utils.staged |
| def _compute_loss_from_registrations( |
| self, |
| func_args: FuncArgsVariants |
| ) -> Array: |
|
|
| loss = self._estimator.compute_func_from_registered( |
| func_args, self._batch_size_extractor(func_args[-1])) |
|
|
| if self._l2_reg > 0.0: |
|
|
| l2_reg_val = self._l2_reg / 2 * utils.squared_norm( |
| func_args[self._params_index]) |
|
|
| loss += l2_reg_val |
|
|
| return loss |
|
|
| @utils.staged |
| def _init( |
| self, |
| params: Params, |
| rng: PRNGKey, |
| batch: Batch, |
| func_state: FuncState | None = None, |
| ) -> State: |
| """A staged function to initialize the optimizer state .""" |
|
|
| |
| |
|
|
| return Optimizer.State( |
| velocities=jax.tree_util.tree_map(jnp.zeros_like, params), |
| estimator_state=self._estimator.init( |
| rng=rng, |
| func_args=make_func_args( |
| params=params, |
| func_state=func_state, |
| rng=rng, |
| batch=self._batch_process_func(batch), |
| has_state=self._value_func_has_state, |
| has_rng=self._value_func_has_rng, |
| ), |
| exact_powers_to_cache=self._exact_powers_to_cache, |
| approx_powers_to_cache=self._approx_powers_to_cache, |
| cache_eigenvalues=False |
| ), |
| damping=jnp.array( |
| (self._initial_damping if self._initial_damping is not None |
| else -1e10), dtype=float), |
| data_seen=jnp.array(0, dtype=int), |
| step_counter=jnp.array(0, dtype=int) |
| ) |
|
|
| def init( |
| self, |
| params: Params, |
| rng: PRNGKey, |
| batch: Batch, |
| func_state: FuncState | None = None, |
| ) -> State: |
| """Initializes the optimizer and returns the appropriate optimizer state. |
| |
| NOTE: please do not jit/pmap or otherwise compile this function with JAX, |
| as this can lead to errors. Compilation is handled internally by the |
| optimizer. |
| |
| NOTE: when ``multi_device`` is ``True``, all of the JAX array arguments to |
| this function (including arrays inside of trees), should have an extra |
| leading axis the size of the number of local devices. |
| |
| Args: |
| params: Example models parameters (used for tracing and shape info). |
| rng: A Jax PRNG key. Unlike the ``rng`` in the step function, should be |
| the same for each host and for each slice in the leading axis (i.e. |
| corresponding to devices) when ``multi_device`` is ``True``. |
| batch: An example batch of the same size as the one passed to ``step`` |
| (or returned from the ``data_iterator``). Used for tracing and shape |
| info. |
| func_state: Example function state (used for tracing and shape info). |
| |
| Returns: |
| The initialized optimizer state. |
| """ |
|
|
| if not self.finalized: |
| self.finalize(params, rng, batch, func_state) |
|
|
| |
| _ = self._maybe_mask_out_unregularized_parameters(params, log_paths=True) |
|
|
| return self._init(params, rng, batch, func_state) |
|
|
| @functools.partial(utils.staged, donate_argnums=[1, 3, 5]) |
| def _burnin( |
| self, |
| params: Params, |
| state: State, |
| rng: Array, |
| batch: Batch, |
| func_state: FuncState | None, |
| damping: Array | None, |
| accumulator: utils.MultiChunkAccumulator, |
| sync: Array | bool, |
| ) -> tuple[State, utils.MultiChunkAccumulator]: |
| """A single burnin step, updating only the curvature estimate.""" |
|
|
| _, _, _, precon_damping = self._setup_state_and_schedules( |
| None, None, |
| state.damping if self._use_adaptive_damping else damping, |
| state.step_counter, state.data_seen) |
|
|
| |
| accumulator = accumulator.copy() |
|
|
| func_args, rng = self._setup_func_args_and_rng( |
| params, rng, batch, func_state) |
|
|
| |
| state.estimator_state = self._update_estimator_curvature( |
| state.estimator_state, |
| func_args, |
| rng, |
| ema_old=1.0, |
| ema_new=1.0, |
| precon_damping=precon_damping, |
| sync=sync, |
| ) |
|
|
| |
| if func_state is not None: |
| out, _ = self._value_and_grad_func(*func_args) |
| _, func_state, _ = extract_func_outputs( |
| out, self._value_func_has_aux, self._value_func_has_state) |
|
|
| accumulator.add(func_state) |
|
|
| return state, accumulator |
|
|
| def _burnin_phase( |
| self, |
| num_steps: int, |
| params: Params, |
| state: State, |
| rng: PRNGKey, |
| data_iterator: Iterator[Batch], |
| func_state: FuncState | None = None, |
| damping: Array | None = None, |
| ) -> tuple[State, FuncState | None]: |
| """Runs all burnin steps required.""" |
|
|
| if num_steps > 0: |
|
|
| rng = self._rng_split(rng, num_steps) |
|
|
| accumulator = utils.MultiChunkAccumulator.zeros_like( |
| func_state, self.multi_device) |
|
|
| for i, rng_i in enumerate(rng): |
| batch = next(data_iterator) |
|
|
| state, accumulator = self._burnin( |
| params, state, rng_i, batch, func_state, damping, accumulator, |
| i == num_steps - 1) |
|
|
| func_state = accumulator.value_and_clear() |
|
|
| return state, func_state |
|
|
| @functools.partial( |
| utils.staged, donate_argnums=(0, 1, 4), static_argnums=(8, 9)) |
| @utils.auto_scope_method |
| def _step( |
| self, |
| params: Params, |
| state: State, |
| rng: Array, |
| batch: Batch, |
| func_state: FuncState | None, |
| learning_rate: Array | None, |
| momentum: Array | None, |
| damping: Array | None, |
| should_update_estimate_curvature: bool, |
| should_update_damping: bool, |
| curvature_ema: Numeric, |
| inverse_update_period: Numeric, |
| )-> ReturnEither: |
| """A single full step of the optimizer.""" |
|
|
| |
| state = state.copy() |
|
|
| |
| (learning_rate, momentum, damping, |
| precon_damping) = self._setup_state_and_schedules( |
| learning_rate, momentum, |
| state.damping if self._use_adaptive_damping else damping, |
| state.step_counter, state.data_seen) |
|
|
| func_args, rng = self._setup_func_args_and_rng( |
| params, rng, batch, func_state) |
|
|
| |
| if should_update_estimate_curvature: |
| if self._share_curvature_and_grad_forward: |
| ( |
| state.estimator_state, |
| loss, |
| grads, |
| ) = self._update_estimator_curvature_and_value_and_grad( |
| state.estimator_state, |
| func_args, |
| rng, |
| ema_old=curvature_ema, |
| ema_new=1.0, |
| precon_damping=precon_damping, |
| sync=self._should_sync_estimator(state, inverse_update_period), |
| ) |
| else: |
| state.estimator_state = self._update_estimator_curvature( |
| state.estimator_state, |
| func_args, |
| rng, |
| ema_old=curvature_ema, |
| ema_new=1.0, |
| precon_damping=precon_damping, |
| sync=self._should_sync_estimator(state, inverse_update_period), |
| ) |
|
|
| del rng |
|
|
| |
| if ( |
| should_update_estimate_curvature |
| and self._share_curvature_and_grad_forward |
| ): |
| func_state = None |
| aux = None |
| else: |
| loss, grads, func_state, aux = self._compute_loss_and_grads( |
| func_args, state=state) |
|
|
| |
| loss, grads = utils.pmean_if_pmap((loss, grads), self.pmap_axis_name) |
|
|
| |
| state = self._maybe_update_inverse_cache( |
| state, precon_damping, inverse_update_period) |
|
|
| |
| preconditioned_gradient = self._compute_preconditioned_gradient( |
| state, grads, precon_damping |
| ) |
|
|
| |
| |
| if self._step_stats_hook is not None: |
| hook_stats = self._step_stats_hook( |
| self._estimator, grads, preconditioned_gradient) |
| else: |
| hook_stats = {} |
|
|
| |
| preconditioned_gradient, scaled_grad_norm_sq = ( |
| self._maybe_apply_norm_constraint( |
| grads, preconditioned_gradient, learning_rate, |
| ) |
| ) |
|
|
| vectors = (preconditioned_gradient, state.velocities) |
|
|
| |
| coefficients, quad_model_change = self._coefficients_and_quad_change( |
| state=state, |
| vectors=vectors, |
| grads=grads, |
| learning_rate=learning_rate, |
| momentum=momentum, |
| damping=damping, |
| func_args=func_args, |
| should_update_damping=should_update_damping, |
| ) |
|
|
| |
| delta = self._weighted_sum_of_objects(vectors, coefficients) |
|
|
| |
| new_params = jax.tree_util.tree_map(jnp.add, params, delta) |
|
|
| if should_update_damping or self._use_step_rejection: |
|
|
| new_loss = self._compute_loss_value((new_params,) + func_args[1:]) |
| |
| new_loss = utils.pmean_if_pmap(new_loss, self.pmap_axis_name) |
|
|
| else: |
| new_loss = self._invalid_metric_value |
|
|
| |
| if should_update_damping: |
|
|
| state.damping, rho = self._compute_new_damping_and_rho( |
| loss, new_loss, quad_model_change, state.damping) |
|
|
| else: |
| |
| |
| new_loss, rho = self._invalid_metric_value, self._invalid_metric_value |
|
|
| if self._use_step_rejection: |
|
|
| reject_step = jnp.logical_or(jnp.isnan(new_loss), new_loss > loss) |
|
|
| params, state.velocities, state.damping = lax.cond( |
| reject_step, |
| lambda: (params, state.velocities, |
| self._reject_damping_increase_factor * state.damping), |
| lambda: (new_params, delta, state.damping)) |
|
|
| else: |
| |
| reject_step = False |
| params, state.velocities = new_params, delta |
|
|
| |
| batch_size = self._batch_size_extractor(func_args[-1]) |
|
|
| if self.multi_device: |
| total_batch_size = batch_size * jax.device_count() |
| else: |
| total_batch_size = batch_size |
|
|
| |
| state.data_seen = state.data_seen + total_batch_size |
| state.step_counter = state.step_counter + 1 |
|
|
| |
| |
| |
| |
| |
| stats = dict( |
| step=state.step_counter, |
| batch_size=jnp.asarray(total_batch_size, dtype=jnp.int32), |
| data_seen=state.data_seen, |
| loss=loss, |
| new_loss=new_loss, |
| learning_rate=-coefficients[0], |
| momentum=coefficients[1], |
| damping=damping, |
| precon_damping=precon_damping, |
| rho=rho, |
| quad_model_change=quad_model_change, |
| scaled_grad_norm_sq=scaled_grad_norm_sq, |
| ) |
|
|
| if self._use_step_rejection: |
| stats["step_rejected"] = reject_step |
|
|
| stats.update(hook_stats) |
|
|
| if aux is not None: |
| aux = utils.pmean_if_pmap(aux, self.pmap_axis_name) |
| stats["aux"] = aux |
|
|
| if self._include_norms_in_stats: |
| stats["param_norm"] = utils.norm(params) |
| stats["grad_norm"] = utils.norm(grads) |
| stats["precon_grad_norm"] = utils.norm(preconditioned_gradient) |
| stats["update_norm"] = utils.norm(delta) |
|
|
| if self._include_per_param_norms_in_stats: |
| stats.update(utils.per_parameter_norm(params, "param_norm")) |
| stats.update(utils.per_parameter_norm(grads, "grad_norm")) |
| stats.update( |
| utils.per_parameter_norm(preconditioned_gradient, "precon_grad_norm") |
| ) |
| stats.update(utils.per_parameter_norm(delta, "update_norm")) |
|
|
| if self._include_registered_loss_in_stats: |
| assert aux is not None |
| stats["loss_registered"] = aux.pop("loss_registered") |
| stats["loss_registered"] = utils.pmean_if_pmap(stats["loss_registered"], |
| self.pmap_axis_name) |
| stats["loss_registered_reldiff"] = ( |
| stats["loss_registered"] - loss) / loss |
|
|
| if self._value_func_has_state: |
| return params, state, func_state, stats |
|
|
| assert func_state is None |
|
|
| return params, state, stats |
|
|
| def step( |
| self, |
| params: Params, |
| state: State, |
| rng: PRNGKey, |
| data_iterator: Iterator[Batch] | None = None, |
| batch: Batch | None = None, |
| func_state: FuncState | None = None, |
| learning_rate: Array | None = None, |
| momentum: Array | None = None, |
| damping: Array | None = None, |
| global_step_int: int | None = None |
| )-> ReturnEither: |
| """Performs a single update step using the optimizer. |
| |
| NOTE: please do not jit/pmap or otherwise compile this function with JAX, |
| as this can lead to errors. Compilation is handled internally by the |
| optimizer. |
| |
| NOTE: when ``multi_device`` is ``True``, all of the JAX array arguments to |
| this function (including arrays inside of trees), should have an extra |
| leading axis the size of the number of local devices. Slices of ``batch`` |
| and ``rng`` should be different for each device, whereas the other arugments |
| should be identical for each slice. Passing the arguments any other way will |
| result in an exception, or possibly undefined behavior. |
| |
| Args: |
| params: The current parameters of the model. |
| state: The current state of the optimizer. |
| rng: A Jax PRNG key. Should be different for each iteration, each host, |
| and for each slice in the leading axis (i.e. corresponding to devices) |
| when ``multi_device`` is ``True``. |
| data_iterator: A data iterator to use (if not passing ``batch``). |
| batch: A single batch used to compute the update. Should only pass one |
| of ``data_iterator`` or ``batch``. |
| func_state: Any function state that gets passed in and returned. |
| learning_rate: Learning rate to use if the optimizer was created with |
| ``use_adaptive_learning_rate=False`` and |
| ``learning_rate_schedule=None``. Should be ``None`` otherwise. |
| momentum: Momentum to use if the optimizer was created with |
| ``use_adaptive_momentum=False`` and ``momentum_schedule=None``. Should |
| be ``None`` otherwise. |
| damping: Damping to use if the optimizer was created with |
| ``use_adaptive_damping=False`` and ``damping_schedule=None``. Should be |
| ``None`` otherwise. See discussion of constructor argument |
| ``initial_damping`` for more information about damping. |
| global_step_int: The global step as a python int. Note that this must |
| match the step internal to the optimizer that is part of its state. |
| |
| Returns: |
| (params, state, stats) if ``value_func_has_state=False`` and |
| (params, state, func_state, stats) otherwise, where |
| |
| * params is the updated model parameters. |
| |
| * state is the updated optimizer state. |
| |
| * func_state is the updated function state. |
| |
| * stats is a dictionary of useful statistics including the loss. |
| """ |
|
|
| if (data_iterator is None) == (batch is None): |
| raise ValueError("Exactly one of the arguments ``data_iterator`` and " |
| "``batch`` must be provided.") |
|
|
| step_counter_int = self._verify_args_and_get_step_counter( |
| step_counter=state.step_counter, |
| learning_rate=learning_rate, |
| momentum=momentum, |
| damping=damping, |
| global_step_int=global_step_int, |
| ) |
|
|
| if step_counter_int == 0: |
|
|
| if self._num_burnin_steps > 0: |
|
|
| if data_iterator is None: |
| raise ValueError("If num_burnin_steps > 0, data_iterator must be " |
| "provided.") |
|
|
| rng, burnin_rng = self._rng_split(rng, 2) |
|
|
| state, func_state = self._burnin_phase( |
| num_steps=self._num_burnin_steps, |
| params=params, |
| state=state, |
| rng=burnin_rng, |
| data_iterator=data_iterator, |
| func_state=func_state, |
| damping=damping, |
| ) |
|
|
| if data_iterator is not None: |
| batch = next(data_iterator) |
|
|
| if (step_counter_int == 0 and self._use_adaptive_damping |
| and self._use_initial_damping_calibration): |
|
|
| assert self._num_burnin_steps > 0 |
|
|
| state = self._calibrate_initial_damping( |
| params, state, rng, batch, func_state, learning_rate, momentum) |
|
|
| should_update_estimate_curvature = self._should_update_estimate_curvature( |
| step_counter_int |
| ) |
| should_update_damping = self._should_update_damping(step_counter_int) |
|
|
| curvature_ema, inverse_update_period = self._live_step_scalars() |
|
|
| return self._step( |
| params, state, rng, batch, func_state, learning_rate, momentum, damping, |
| should_update_estimate_curvature, should_update_damping, |
| curvature_ema, inverse_update_period) |
|
|
| def _calibrate_initial_damping( |
| self, |
| params: Params, |
| state: State, |
| rng: PRNGKey, |
| batch: Batch, |
| func_state: FuncState | None = None, |
| learning_rate: Array | None = None, |
| momentum: Array | None = None, |
| ) -> State: |
| """Calibrates the initial damping parameter.""" |
|
|
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
|
|
| while True: |
|
|
| prev_damping = float(self.get_first(state.damping)) |
|
|
| |
| |
| |
| curvature_ema, inverse_update_period = self._live_step_scalars() |
|
|
| ret = self._step( |
| self.copy_obj(params), self.copy_obj(state), rng, batch, |
| self.copy_obj(func_state), learning_rate, momentum, None, False, True, |
| curvature_ema, inverse_update_period) |
|
|
| new_state = ret[1] |
|
|
| new_damping = float(self.get_first(new_state.damping)) |
| state.damping = new_state.damping |
|
|
| del new_state |
|
|
| if prev_damping == new_damping: |
| return state |
|
|
| @utils.auto_scope_method |
| def _compute_exact_quad_model_filtered( |
| self, |
| vectors: Sequence[Params], |
| grads: Params, |
| func_args: FuncArgsVariants, |
| state: State | None = None, |
| fixed_coefficients: Sequence[Numeric | None] | None = None, |
| **kwargs, |
| ) -> QuadModelParams: |
| """Computes the components of the exact quadratic model.""" |
|
|
| |
| |
| |
| |
|
|
| if fixed_coefficients is None: |
| return self._compute_exact_quad_model( |
| vectors, grads, func_args, state=state, **kwargs) |
|
|
| assert len(vectors) == len(fixed_coefficients) |
| assert len(vectors) == 2 |
|
|
| def if_momentum_coeff_zero(): |
|
|
| |
| quad_model = self._compute_exact_quad_model( |
| vectors[:1], grads, func_args, state=state, **kwargs) |
|
|
| |
| return tuple( |
| jnp.pad(arr, [(0, 1)] * arr.ndim, constant_values=0.0) |
| for arr in quad_model |
| ) |
|
|
| |
| if (isinstance(fixed_coefficients[1], float) |
| and fixed_coefficients[1] == 0.0): |
|
|
| return if_momentum_coeff_zero() |
|
|
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
|
|
| return self._compute_exact_quad_model( |
| vectors, grads, func_args, state=state, **kwargs) |
|
|
| def _maybe_mask_out_unregularized_parameters( |
| self, params: Params, log_paths: bool = False) -> Params: |
| """Mask out parameters that are not l2 regularized.""" |
|
|
| if log_paths: |
| logging.info("Unregularized parameters masking info (for curvature " |
| "calculations and L2 regularization)") |
|
|
| def maybe_mask_out_single_param( |
| path: tuple[Any, ...], |
| param: Array |
| ) -> Array: |
| """Zero out a single parameter.""" |
| str_path = [] |
| for p in path: |
| if isinstance(p, jax.tree_util.DictKey): |
| str_path.append(p.key) |
| elif isinstance(p, jax.tree_util.GetAttrKey): |
| str_path.append(p.name) |
|
|
| should_mask = any( |
| p in str_path |
| for p in self._regularized_parameters_path_exclusions |
| ) |
|
|
| if log_paths: |
| log_message = "Masking" if should_mask else "Not masking" |
| logging.info(" %s out %s", log_message, path) |
|
|
| return jnp.zeros_like(param) if should_mask else param |
|
|
| return jax.tree.map_with_path( |
| maybe_mask_out_single_param, params |
| ) |
|
|
| @utils.auto_scope_method |
| def _compute_exact_quad_model( |
| self, |
| vectors: Sequence[Params], |
| grads: Params, |
| func_args: FuncArgsVariants, |
| state: State | None = None, |
| ) -> QuadModelParams: |
| """Computes the components of the exact quadratic model. |
| |
| See comments of QuadModelParams for a description of the returned tuple. |
| |
| Args: |
| vectors: sequence of update vectors `V`. |
| grads: The gradient `g` of the loss function. |
| func_args: The arguments to the model's value function. |
| state: The current optimizer state. |
| |
| Returns: |
| A `QuadModelParams` tuple (A, D, R, b). |
| """ |
|
|
| del state |
|
|
| if self._mat_type_for_exact_quad_model == "fisher": |
| c_factor_v = tuple(self._implicit.multiply_fisher_factor_transpose |
| (func_args, vi) for vi in vectors) |
| elif self._mat_type_for_exact_quad_model == "ggn": |
| c_factor_v = tuple(self._implicit.multiply_ggn_factor_transpose |
| (func_args, vi) for vi in vectors) |
| else: |
| raise ValueError(f"Unrecognized matrix type string for exact quad model:" |
| f"'{self._mat_type_for_exact_quad_model}'.") |
|
|
| masked_vectors = tuple(self._maybe_mask_out_unregularized_parameters(vi) |
| for vi in vectors) |
|
|
| |
| A = utils.matrix_of_inner_products(c_factor_v) |
| D = utils.matrix_of_inner_products(vectors) |
| R = utils.matrix_of_inner_products(masked_vectors) |
| b = utils.vector_of_inner_products(grads, vectors) |
| |
|
|
| quad_model_params = (A, D, R, b) |
|
|
| return utils.pmean_if_pmap(quad_model_params, self.pmap_axis_name) |
|
|
| @functools.partial(utils.staged, donate_argnums=2) |
| @utils.auto_scope_method |
| def _compute_approx_quad_model( |
| self, |
| state: State, |
| vectors: Sequence[Params], |
| grads: Params, |
| ) -> QuadModelParams: |
| """Computes the components of the approximate quadratic model.""" |
|
|
| |
| def c_times_v(v): |
| return self._estimator.multiply( |
| state=state.estimator_state, |
| parameter_structured_vector=v, |
| identity_weight=0.0, |
| exact_power=True, |
| use_cached=False, |
| pmap_axis_name=self.pmap_axis_name, |
| norm_to_scale_identity_weight_per_block=self._norm_to_scale_identity_weight_per_block, |
| ) |
|
|
| c_vectors = [c_times_v(v_i) for v_i in vectors] |
|
|
| return (utils.symmetric_matrix_inner_products(c_vectors, vectors), |
| utils.matrix_of_inner_products(vectors), |
| utils.matrix_of_inner_products(vectors), |
| utils.vector_of_inner_products(grads, vectors)) |
|
|
| def _evaluate_quadratic_model( |
| self, |
| a: Array, |
| a_damped: Array, |
| b: Array, |
| w: Array, |
| ) -> Array: |
| """Computes the quadratic model value from the inputs provided.""" |
|
|
| a_final = a_damped if self._include_damping_in_quad_change else a |
|
|
| return jnp.dot(w, jnp.dot(a_final, w)) / 2 + jnp.dot(w, b) |
|
|
| @utils.staged |
| def _solve_quad_model( |
| self, |
| quad_model_parameters: QuadModelParams, |
| damping: Array, |
| fixed_coefficients: Sequence[Numeric | None], |
| reg_coeff: Numeric | None = None, |
| ) -> tuple[tuple[Numeric, ...], Array]: |
| """Solves for the optimal learning rate and momentum of the quadratic model. |
| |
| Args: |
| quad_model_parameters: The computed matrices A, D, R, and vector b. |
| damping: The damping to use for evaluating the quadratic model. |
| fixed_coefficients: A list over the vectors of the fixed numerical values |
| to use for their coefficients. For each of these that is None, the |
| quadratic model is minimized to compute the 'optimal' coefficient value. |
| reg_coeff: The L2 regularization parameter to use. If None, the default |
| value from the optimizer is used. |
| |
| Returns: |
| A tuple of coefficients which are the solution (and include any values that |
| are not None from fixed_weights), and the value of the quadratic model |
| function for this solution (as a scalar). |
| |
| Raises: |
| The function currently supports only up to two vectors, hence if you |
| provide more, it will raise a ``NotImplementedError``. |
| """ |
|
|
| if reg_coeff is None: |
| |
| reg_coeff = self._l2_reg |
|
|
| |
| A_no_diag, D, R, b = quad_model_parameters |
| A = A_no_diag + reg_coeff * R |
| A_damped = A + damping * D |
|
|
| if all(c is None for c in fixed_coefficients): |
| |
|
|
| if len(fixed_coefficients) == 1: |
| |
| |
| special_case = jnp.logical_and(A_damped[0, 0] == 0, b[0] == 0) |
| w = -lax.cond(special_case, lambda: b, lambda: b / A_damped[0]) |
|
|
| elif len(fixed_coefficients) == 2: |
| w = -utils.psd_solve_maybe_zero_last_idx(A_damped, b) |
|
|
| else: |
| raise NotImplementedError() |
|
|
| elif all(c is not None for c in fixed_coefficients): |
| |
|
|
| w = jnp.asarray(fixed_coefficients) |
|
|
| elif len(fixed_coefficients) == 2: |
| |
|
|
| w = [None, None] |
| index = fixed_coefficients.index(None) |
| w[1 - index] = fixed_coefficients[1 - index] |
|
|
| b_extra = A_damped[1 - index, index] * w[1 - index] |
| |
|
|
| w[index] = -(b[index] + b_extra) / A_damped[index, index] |
|
|
| else: |
| raise NotImplementedError() |
|
|
| w = tuple(w) |
| w: tuple[Numeric, ...] |
|
|
| quad_model_change = self._evaluate_quadratic_model( |
| A, A_damped, b, jnp.array(w)) |
|
|
| return w, quad_model_change |
|
|
| @utils.staged |
| def _compute_new_damping_and_rho( |
| self, |
| old_loss: Array, |
| new_loss: Array, |
| quad_change: Array, |
| current_damping: Array, |
| ) -> tuple[Array, Array]: |
| """Computes the reduction ratio and the updated value of the damping.""" |
|
|
| |
| rho = (new_loss - old_loss) / quad_change |
| rho_not_nan = jnp.nan_to_num(rho, nan=-100.0) |
|
|
| |
| should_increase = rho_not_nan < self._damping_lower_threshold |
| increased_damping = current_damping / self._damping_decay_factor |
| should_decrease = rho_not_nan > self._damping_upper_threshold |
| decreased_damping = current_damping * self._damping_decay_factor |
|
|
| damping = jnp.select([should_decrease, should_increase], |
| [decreased_damping, increased_damping], |
| default=current_damping) |
|
|
| return jnp.clip(damping, self._min_damping, self._max_damping), rho |
|
|
| @utils.staged |
| def _weighted_sum_of_objects( |
| self, |
| objects: Sequence[utils.PyTree], |
| coefficients: Sequence[Numeric], |
| ) -> utils.PyTree: |
| """Returns the weighted sum of the objects in the sequence.""" |
| return utils.weighted_sum_of_objects(objects, coefficients) |
|
|
|
|
| def convert_value_and_grad_to_value_func( |
| value_and_grad_func: ValueAndGradFunc, |
| has_aux: bool = False, |
| ) -> ValueFunc: |
| """Converts a value_and_grad function to value_func only. |
| |
| Args: |
| value_and_grad_func: The function which computes the loss value and the |
| gradients w.r.t. parameters. |
| has_aux: Similar to the meaning in :func:`jax.grad`, whether the |
| ``value_and_grad_func`` returns with the loss value any auxiliary data. |
| |
| Returns: |
| A function that returns only the loss value. |
| """ |
|
|
| def value_func(*args, **kwargs) -> Array: |
| out, _ = value_and_grad_func(*args, **kwargs) |
| return out[0] if has_aux else out |
|
|
| return value_func |
|
|
|
|
| def convert_value_and_grad_to_clean_value_and_grad( |
| value_and_grad_func: ValueAndGradFunc, |
| has_aux: bool = False, |
| ) -> utils.ValueAndGradFunc: |
| """Converts a value_and_grad function to return only (loss, grads). |
| |
| Args: |
| value_and_grad_func: The function which computes the loss value and the |
| gradients w.r.t. parameters. |
| has_aux: Similar to the meaning in :func:`jax.grad`, whether the |
| ``value_and_grad_func`` returns with the loss value any auxiliary data. |
| |
| Returns: |
| A function that returns `(loss, grads)`. |
| """ |
|
|
| def clean_value_and_grad_func(*args, **kwargs) -> tuple[Array, Params]: |
| out, grads = value_and_grad_func(*args, **kwargs) |
| loss = out[0] if has_aux else out |
| return loss, grads |
|
|
| return clean_value_and_grad_func |
|
|
|
|
| def make_func_args( |
| params: Params, |
| func_state: FuncState | None, |
| rng: PRNGKey | None, |
| batch: Batch, |
| has_state: bool, |
| has_rng: bool, |
| ) -> FuncArgsVariants: |
| """Constructs the arguments to the model function in the pre-assumed order. |
| |
| The model function is assumed to take arguments in the following order: |
| params, func_state, rng, batch |
| If it has no function state or does not use an rng, those two arguments are |
| discarded. |
| |
| Args: |
| params: The model parameters. |
| func_state: The function state, if ``has_state`` is ``True``, ``None`` |
| otherwise. |
| rng: The PRNG, if ``has_rng`` is ``True``, ``None`` otherwise. |
| batch: The batch of data. |
| has_state: Whether the function has a function state. |
| has_rng: Whether the function uses an rng. |
| |
| Returns: |
| The arguments that need to be passed to the model function. |
| """ |
| if has_state and func_state is None: |
| raise ValueError("`func_state=None`, but argument `has_state=True`.") |
|
|
| if has_rng and rng is None: |
| raise ValueError("`rng=None`, but argument `has_rng=True`.") |
|
|
| if not has_state and not has_rng: |
| return params, batch |
|
|
| elif not has_rng: |
| return params, func_state, batch |
|
|
| elif not has_state: |
| return params, rng, batch |
|
|
| else: |
| return params, func_state, rng, batch |
|
|
|
|
| def extract_func_outputs( |
| raw_outputs: FuncOutputs, |
| has_aux: bool, |
| has_state: bool, |
| ) -> tuple[Array, FuncState | None, FuncAux | None]: |
| """Converts the raw output of the model function into loss,func_state and aux. |
| |
| Args: |
| raw_outputs: The direct output of the model function. |
| has_aux: Whether the model function returns also some auxiliary data. |
| has_state: Whether the model function has a function state. |
| |
| Returns: |
| A triple ``(loss, func_state, aux)``. If the model function does not return |
| any auxiliary data than ``aux`` will be ``None`` and if it does not have a |
| state ``func_state`` will be ``None``. |
| """ |
|
|
| if not has_aux and not has_state: |
| assert isinstance(raw_outputs, Array) |
| return raw_outputs, None, None |
|
|
| loss, other = raw_outputs |
|
|
| if has_aux and has_state: |
| func_state, aux = other |
| elif has_aux: |
| func_state, aux = None, other |
| else: |
| func_state, aux = other, None |
|
|
| return loss, func_state, aux |
|
|