skip to content

Rescaling a ReLU network leaves its function identical but inflates measured sharpness — what does that break?

level: principalimportance: nice to knowfreq 18%

answer

  1. the network computes exactly the same thing
  2. ReLU is positively homogeneous
  3. scale one layer up, the next down
  4. eigenvalues change, predictions do not
  5. sharpness is not reparameterization-invariant

basics

~10 s

It breaks sharpness as a standalone explanation of generalization. Scaling one layer up and the next down computes the same function with the same test error, yet raw curvature measures change freely.

solid answer

~40 s

A rectified linear unit is positively homogeneous: for c > 0, ReLU(c*z) = c*ReLU(z). So you can multiply one layer's weights and biases by c and divide the next layer's weights by c, and the network computes exactly the same function: same predictions, same training loss, same test error. But the curvature at that point changes, because the layer whose weights shrank now sees the loss climb far faster per unit of absolute perturbation, so Hessian eigenvalues and fixed-radius sharpness inflate without limit. Raw sharpness is therefore a property of the parameterization, not of the learned function, and cannot by itself explain generalization. The fix is a measure invariant to that rescaling: perturb each coordinate proportionally to its own magnitude.

code

python · 21 lines
python
import random

def net(x, w1, b1, w2):
    return w2 * max(0.0, w1 * x + b1)          # one ReLU hidden unit

w1, b1, w2 = 0.7, 0.1, 1.3
c = 10.0
r1, rb, r2 = w1 * c, b1 * c, w2 / c            # identical function, rescaled

xs = [0.3, 1.0, 2.5]
print([abs(net(x, w1, b1, w2) - net(x, r1, rb, r2)) < 1e-9 for x in xs])

def loss(p):                                   # distance from the original outputs
    return sum((net(x, *p) - net(x, w1, b1, w2)) ** 2 for x in xs)

def sharpness(p, radius=0.05, trials=2000):    # worst rise in a fixed-radius ball
    random.seed(0)
    return max(loss([v + random.uniform(-radius, radius) for v in p])
               for _ in range(trials))

print(round(sharpness((w1, b1, w2)), 4), round(sharpness((r1, rb, r2)), 4))

go deeper

for a junior

Know the fact rather than the theory: two very different-looking weight sets can compute exactly the same function, so a number read off the weights need not say anything about the model's behaviour.

for a middle

Be able to derive the construction. State the positive-homogeneity identity for ReLU, show which weights and biases get multiplied and divided, and explain why a fixed absolute perturbation then hits the shrunken layer much harder.

for a senior

Show the operating consequence: only compare sharpness between checkpoints from the same recipe with comparable weight norms, prefer a magnitude-proportional perturbation, and report a radius sweep rather than a single number.

for a principal

Take a position on when a curvature statistic may influence a decision at all. Argue for invariance as a precondition, keep held-out evaluation as the actual gate, and push back on both the team that ships on a sharpness score and the team that dismisses flatness entirely.

## The construction The rectified linear unit satisfies ReLU(c*z) = c*ReLU(z) for any c > 0: it is positively homogeneous of degree one. Take two consecutive layers in a ReLU network, with the first computing h = ReLU(W1 x + b1) and the second computing W2 h. Now form a rescaled parameter set: W1' = c * W1, b1' = c * b1, W2' = W2 / c The hidden activation becomes ReLU(c*W1 x + c*b1) = c * ReLU(W1 x + b1) = c*h, and the second layer multiplies it by W2/c, returning exactly W2 h. Every output for every input is bit-for-bit the same function. Training loss, validation loss, test accuracy, calibration — all unchanged, because the map from inputs to outputs is unchanged. ## Why the sharpness number moves anyway Sharpness is measured in parameter space, and the two parameter vectors are far apart even though the functions coincide. Consider a fixed-radius probe that adds a perturbation of absolute size delta to every weight. In the rescaled network the second layer's weights are c times smaller, so a perturbation of size delta is a *relatively* enormous change to them — and those weights now multiply an activation that is c times larger. The loss therefore climbs much faster per unit of absolute perturbation. Formally, the Hessian is transformed by the reparameterization and its eigenvalues along the affected directions grow; by choosing c large you can make the same solution look arbitrarily sharp. This is not a subtlety to file away. It means that a fixed-radius perturbation measure, the Hessian spectral norm and any other raw curvature statistic are properties of the *coordinates*, not of the function. Whatever they are correlated with in a typical experiment, they cannot be the mechanism, because they can be moved without touching the thing they are supposed to predict. ## What this does and does not refute It does not say flat solutions are unrelated to generalization. Within a fixed parameterization and a fixed training recipe, weight scales are not arbitrary — they are set by initialization, weight decay and the optimizer — so the comparison between two checkpoints from the same pipeline is usually meaningful, and the empirical correlation people observe is real. What it refutes is the strong claim: that sharpness is *the* explanation, and that a curvature number can be read off any two arbitrary models and used to rank them. Two networks from different architectures, different normalization choices or different weight-decay settings can differ in raw sharpness for reasons that have nothing to do with the functions they compute. The same degeneracy appears wherever the loss surface has scale symmetries. When a layer's output passes through a normalization that removes its scale, multiplying that layer's weights by a constant leaves the function unchanged too — so the parameterization is again not identified, and raw curvature in those directions is not meaningful on its own. ## The repair Make the measure invariant to the symmetry. Instead of perturbing every weight by the same absolute amount, perturb each coordinate proportionally to its own current magnitude — an element-wise adaptive radius. Under the rescaling above, the perturbation applied to the shrunken layer shrinks with it, and the measured sharpness stays put. The same idea can be applied to a sharpness-penalizing training objective, whose radius otherwise means different things in different layers and transfers badly between models. Other partial repairs: normalize by the weight norm and report a relative rather than absolute sharpness; restrict comparisons to checkpoints with matched weight norms; or report the whole radius sweep rather than a single number, so a reader can see whether two curves differ in shape or only in scale. ## The judgment call The practical stance a lead should hold is this. Sharpness is a useful *diagnostic within a controlled comparison* — same architecture, same recipe, same weight scale — and a bad *model-selection score across arbitrary models*. If someone proposes gating a release on a curvature number, ask what makes that number invariant, ask whether the candidates have comparable weight norms, and insist on a held-out metric as the actual decision criterion. And if someone dismisses flatness entirely because of the rescaling counterexample, the correct reply is that the counterexample indicts the *measure*, not the intuition: fix the measure, keep the intuition, and never let either replace an evaluation on held-out data.

  • How does an element-wise adaptive radius repair the measure?
    Instead of perturbing every weight by the same absolute amount, scale each coordinate's perturbation by that coordinate's own magnitude. When a layer's weights are divided by c, the perturbation applied to them is divided by c as well, so the measured loss rise is unchanged by the rescaling. The measure becomes invariant to exactly the symmetry that broke the fixed-radius version, and the radius hyperparameter transfers better across layers and models.
  • Does this mean flatness is a useless idea?
    No — it indicts the measure, not the intuition. Within one architecture and one training recipe the weight scales are pinned by initialization, weight decay and the optimizer, so comparing two checkpoints from the same pipeline remains informative. What the counterexample forbids is treating a raw curvature number as a portable score that ranks arbitrary models, or as an explanation of generalization that stands on its own.
  • Where else does this scale degeneracy show up in modern networks?
    Anywhere the loss surface has a scale symmetry. When a layer's output is normalized so that its scale is removed, multiplying that layer's weights by a constant leaves the computed function unchanged, so the parameterization is again unidentified and raw curvature in those directions carries no functional meaning. The scale instead shows up in how large the effective parameter steps are, which is a different conversation from generalization.

Measuring a solution's sharpness with a fixed absolute radius is like judging a country's wealth by cash on hand without saying the currency. Redenominate and the number changes wildly while nothing real has moved.

saying these in an interview costs you the question

  • Claims the rescaling changes the network's predictions
  • Says the trick only works because ReLU is non-differentiable at zero
  • Treats raw Hessian eigenvalues as a portable model-selection score
  • Concludes flatness is meaningless rather than the measure being wrong
  • Forgets that the first layer's bias must be scaled too

context