Skip to content

NNX spectral and weight normalisation writes normalised weights back into params diverging from Linen and from the papers' update rules #5576

Description

@PizzasBear

I looked at the implementation of nnx.SpectralNorm and found out that it modified the parameters in place, which is a real difference from linen.SpectralNorm and the original definition in the paper Spectral Normalization for Generative Adversarial Networks. I searched for another place where this occurs and realised this also applies to WeightNorm which also saves the normalised parameter rather than restoring it. This alters the intended behaviour as specified by the appropriate paper Weight Normalization: A Simple Reparameterization to Accelerate Training of Deep Neural Networks.

To be precise, the weight normalisation paper states:

The idea of normalizing the weight vector has been proposed before (e.g. N. Srebro and A. Shraibman. Rank, trace-norm and max-norm) but earlier work typically still performed optimization in the $\mathbf{w}$-parameterization, only applying the normalization after each step of stochastic gradient descent. This is fundamentally different from our approach: we propose to explicitly reparameterize the model and to perform stochastic gradient descent in the new parameters $\mathbf{v},g$ directly.

Which implies the weight normalisation approach doesn't apply normalisation after each step of SGD, unlike the current NNX implementation which applies it after every forward pass. And furthermore:

Due to projecting away from $\mathbf{w}$, the norm of $\mathbf{v}$ grows monotonically with the number of weight updates when learning a neural network with weight normalization using standard gradient descent without momentum: ...

Which means the parameter norm $\mathbf{v}$ is expected to be allowed to shift freely, while under the current NNX implementation it is pinned to $\lvert g \rvert$ after each forward operation.

The spectral normalisation paper explicitly states the algorithm which doesn't contain pinning the weights to the normalised state.

Algorithm 1 SGD with spectral normalization

  • Initialize $\tilde{\mathbf u}_l\in \mathcal{R}^{d_l}~{\rm for}~l=1,\dots,L$ with a random vector (sampled from isotropic distribution).
  • For each update and each layer $l$:
    • Apply power iteration method to a unnormalized weight $W^l$:
      $\tilde{\mathbf v}_l \leftarrow (W^{l})^{\rm T} \tilde{\mathbf u}_l/|(W^{l})^{\rm T} \tilde{\mathbf u}_l|_2$
      $\tilde{\mathbf u}_l \leftarrow W^{l} \tilde{\mathbf v}_l/|W^l \tilde{\mathbf v}_l|_2$
    • Calculate $\bar{W}_{\rm SN}$ with the spectral norm:
      $\bar{W}_{\rm SN}^l(W^l) = W^l / \sigma(W^l),\ {\rm where}\ \sigma(W^l)=\tilde{\mathbf u}_l^{\rm T} W^l \tilde{\mathbf v}_l$
    • Update $W^l$ with SGD on mini-batch dataset $\mathcal{D}_M$ with a learning rate $\alpha$:
      $W^l \leftarrow W^l - \alpha \nabla_{W^l} \ell(\bar{W}_{\rm SN}^l(W^l), \mathcal{D}_M)$

To fix the issue, we would need to preserve parameter values like how the Linen API did it. Specifically, a simple way to do it is to restore the modified parameters back to their old values before exiting __call__(...).

def __call__(self, x: Array, ...) -> Array:
  # ...

  state = nnx.state(self.layer_instance) # or nnx.state(self.layer_instance, nnx.Param)
  originals = []                                    # new!
  for path, param in nnx.to_flat_state(state):
    originals.append((param, param[...]))           # new!

    self._weightnorm_inplace(path, param)
    # or
    self._spectral_normalize_inplace(path, param, update_stats=update_stats)

  try:                                              # new!
    return self.layer_instance(x, ...)  # type: ignore
  finally:                                          # new!
    for param, original_value in originals:         # new!
      param[...] = original_value                   # new!

One thing to note is that fixing this will introduce a breaking change. Furthermore it breaks the current example in the nnx.WeightNorm doc:

>>> import jax
>>> import numpy as np
>>> from flax import nnx

>>> class Foo(nnx.Module):
...   def __init__(self, rngs: nnx.Rngs):
...     self.normed_linear = nnx.WeightNorm(
...       nnx.Linear(8, 4, rngs=rngs),
...       variable_filter=nnx.PathContains('kernel'),
...       rngs=rngs,
...     )
...
...   def __call__(self, x: jax.Array) -> jax.Array:
...     return self.normed_linear(x)

>>> rng = jax.random.key(42)
>>> model = Foo(rngs=nnx.Rngs(rng))

>>> x = jax.random.normal(rng, (5, 8))
>>> y = model(x)
>>> y.shape
(5, 4)

>>> w = model.normed_linear.layer_instance.kernel[...]
>>> col_norms = np.linalg.norm(np.array(w), axis=0)
>>> np.testing.assert_allclose(col_norms, np.ones(4)) 
### ^--- Applying this fix would break this assertion!

Activity

  1. mohsinm-dev commented on Aug 30, 2026

    @mohsinm-dev
    Contributor

    I think this is valid

    Linen applies the normalized weight only for the wrapped forward while NNX currently writes it back into the underlying Param. So the optimizer ends up updating from the normalized value instead of the original W / v that the gradient was computed with respect to. For WeightNorm this also removes the radial degree of freedom of v that the paper explicitly relies on

    save/restore looks like a reasonable fix here. I would just make sure the normalization itself is also inside the try/finally, so if normalization of one of the params fails we still restore the params already modified

    also probably better to use

    state = nnx.state(self.layer_instance, nnx.Param)

    in WeightNorm.__call__ as well

    I think we should add an optimizer step equivalence test with Linen in addition to the current forward equivalence tests. The current tests can pass even though the persistent parameter state after the forward is different

    one separate thing I noticed while checking this is that WeightNorm.scales are stored as nnx.data rather than nnx.Param, so g does not seem to be part of the default nnx.Param grad/optimizer path. Probably worth handling separately

  2. PizzasBear commented on Aug 30, 2026

    @PizzasBear
    Author

    Actually, I also noticed the same problem with nnx.WeightNorm.scales and I have already opened a separate issue for it: #5577.

  3. samanklesaria commented on Sep 14, 2026

    @samanklesaria
    Collaborator

    Originally, I had thought this wouldn't make a difference. But after some thought, I think this an important difference after all. In the case where we don't mutate in place, let our non-normalized weight matrix after $t$ steps of gradient descent be $W_t$. In the case where we do, let the normalized weight matrix after $t$ steps of gradient descent be $V_t$. I claim that $V_t$ is just the normalized version of $W_t$ for all $t$, but only if you divide your step size for $V_t$ updates by $c_t^2$ where $c_t$ is product of the normalization factors you've seen so far. Specifically, say $W_t = c_t V_t$. If $f$ uses normalization, $f(cW) = f(W)$, so $\nabla f(cW) = c^{-1} \nabla f(W)$. Then $W_{t+1} = c_tV_t - \frac{\eta}{c_t} \nabla f(V_t) = c_t(V_t - \frac{\eta}{c_t^2}\nabla f(V_t)) = c_t \sigma(B_t) V_{t+1}$ where $B_t = V_t - \frac{\eta}{c_t^2}\nabla f(V_t)$. That's just $c_{t+1} V_{t+1}$. So by doing the update in place like this, we end up using different learning rates than we otherwise would.

    Instead of using the _inplace functions and then restoring parameters back to their old values before exiting, I would just manually divide by the norm as in Algorithm 1. This is how equinox does it: https://github.com/patrick-kidger/equinox/blob/3990f4b5b37c9946fa4a9d44b65bf05e5f5b5640/equinox/nn/_spectral_norm.py#L29

    @PizzasBear or @mohsinm-dev do you want to add a PR? Or shall I?

  4. mohsinm-dev commented on Sep 16, 2026

    @mohsinm-dev
    Contributor

    @samanklesaria I can work on it, will open the PR.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions