How does a regression head that outputs a mean and a log-variance train under Gaussian negative log-likelihood?
answer
- Two outputs, one scalar loss
- Residual scaled by inverse predicted variance
- One term punishes claiming large spread
- Exponentiate to keep the variance positive
- Inflating spread is the easy descent direction
basics
~20 sA head emits two numbers per input, a mean and s = log variance, trained with 0.5 * (s + (y - mean)^2 * exp(-s)). The exponential weights each residual by one over the predicted variance; the s term punishes inflating it.
solid answer
~50 sInstead of assuming one fixed noise scale, the head emits `mu(x)` and `s(x) = log(sigma(x)^2)`, and you train both with the Gaussian negative log-likelihood: `L = 0.5 * (s + (y - mu)^2 * exp(-s))`, dropping the constant `0.5 * log(2 * pi)`. Two forces balance. The residual term is multiplied by `exp(-s) = 1/sigma^2`, so examples the model declares noisy contribute a down-weighted mean gradient; the `0.5 * s` term grows with `s`, so declaring everything noisy is punished. Emitting the log rather than the variance keeps the implied variance positive without a constrained output and keeps the exponent stable. For travel-time prediction — quiet motorway segments, chaotic city-centre segments — this lets the objective stop spending its budget fighting irreducible variation. The usual failure is a variance shortcut early in training: raising `s` cuts the loss faster than improving `mu`, the mean gradient is throttled, and the mean fit stalls.
go deeper
Know that a regression head can emit a second number describing how spread out the target is, and that training then uses a likelihood-shaped loss rather than plain squared error.
Be able to write the loss as 0.5 * (s + squared residual * exp(-s)), explain that exp(-s) weights residuals by inverse variance, and say why the head emits a log rather than a raw variance.
Demonstrate that you have trained one: describe the variance shortcut, how you detected it by tracking loss components separately, and which warm-up or clamping strategy you used to get the mean fitted first.
Decide whether the extra head is worth its failure modes at all. Weigh a learned noise scale against simply supplying known per-example weights, and set the expectation that evaluation must change alongside the objective.
## Relaxing the constant-variance assumption Plain squared error is the Gaussian negative log-likelihood with a fixed, input-independent variance. A heteroscedastic head keeps the Gaussian but lets the spread depend on the input: the network emits two scalars per example, a mean `mu(x)` and a raw score `s(x)` interpreted as `log(sigma(x)^2)`. ## The loss Starting from `-log p(y|x)` for `Normal(mu, sigma^2)`: ``` -log p = (y - mu)^2 / (2 * sigma^2) + 0.5 * log(sigma^2) + 0.5 * log(2 * pi) ``` Substituting `sigma^2 = exp(s)` and dropping the constant gives the trainable form: ``` L = 0.5 * ( s + (y - mu)^2 * exp(-s) ) ``` Both outputs receive gradients from this one scalar. Differentiating: ``` dL/dmu = -(y - mu) * exp(-s) dL/ds = 0.5 * ( 1 - (y - mu)^2 * exp(-s) ) ``` Read those two lines carefully, because they contain the whole behaviour of the method. ## What the two gradients do **The mean gradient is inverse-variance weighted.** It is the ordinary squared-error gradient multiplied by `1/sigma^2`. Examples the model believes are precisely predictable pull hard on `mu`; examples it believes are noisy pull weakly. This is the intended benefit: on a travel-time model, a free-flowing motorway segment at 3 a.m. and a city-centre segment at rush hour are not equally learnable, and squared error forces the network to burn capacity chasing residuals in the second that no model could remove. Inverse-variance weighting lets it stop. **The variance gradient is a self-calibrating balance.** `dL/ds` is zero exactly when `exp(s) = (y - mu)^2` — that is, when the predicted variance equals the squared residual. Above that the gradient pushes `s` down, below it pushes `s` up. Averaged over examples sharing an input region, the stationary point is where the predicted variance matches the mean squared residual there. The `0.5 * s` term is what makes this work: without it, `L` would be minimised by sending `s` to infinity, since `exp(-s)` would drive the residual term to zero at no cost. ## Why the log parametrisation Three reasons, all practical: 1. **Positivity for free.** A variance must be positive. Emitting `s` on the whole real line and exponentiating guarantees that without a constrained output or a clamp that would kill gradients at its boundary. 2. **Numerical range.** Real noise scales can span orders of magnitude; a log-scale output covers that range with ordinary-sized activations. 3. **Better-conditioned gradients.** `dL/ds` is bounded below by `0.5` on one side and only grows linearly in the normalised squared residual, whereas parametrising `sigma^2` directly puts it in a denominator, so a near-zero prediction produces a gradient that explodes. It is still wise to keep `s` in a sane interval, because a runaway negative `s` makes `exp(-s)` enormous and reproduces the exploding-gradient problem from the other direction. ## The characteristic failure: the variance shortcut Early in training `mu` is bad everywhere, so residuals are large everywhere. The optimiser has two ways to reduce the loss: improve `mu`, which is slow and requires learning real structure, or raise `s`, which is immediate and requires almost nothing. Raising `s` is often the faster descent direction — and it is self-reinforcing, because raising `s` multiplies the mean gradient by a smaller `exp(-s)`, which makes improving `mu` even slower. The observable symptom is a loss that falls smoothly while the mean fit stays poor, especially in exactly the regions that were hardest at initialisation. Common mitigations, all of which amount to giving the mean a head start: - warm up with plain squared error and switch on the variance term after the mean is roughly fitted; - hold `s` at a constant for the first phase, then unfreeze it; - clamp `s` to a range chosen from the target's overall scale; - alternate: fit the mean, then fit the variance against frozen residuals. ## Reading the trained model correctly Two cautions worth stating. First, the mean produced by this objective is no longer the plain conditional mean of a squared-error fit: it is a precision-weighted fit, so if your evaluation is plain squared error the numbers can look worse even when the model is better calibrated. Second, the `sigma^2` head is a component of a training objective, and what it does and does not license you to say about a prediction's uncertainty is a separate topic with its own subtleties — do not assume the two are the same conversation. ## When it is worth it Use it when the noise scale genuinely varies with the input and you can point to why — different sensors, different regimes, different aggregation levels. Skip it when the target is roughly uniformly noisy, because you have doubled the head, added an optimisation failure mode, and bought nothing.
- What happens to the loss if you drop the log-variance term and keep only the weighted residual?It degenerates. With only `(y - mu)^2 * exp(-s)`, the optimiser drives `s` upward without limit: every extra unit of `s` shrinks the residual term geometrically at zero cost, so the loss goes to zero with an arbitrarily bad mean. The `0.5 * s` term is precisely the price of claiming to be uncertain, and it is what makes the stationary point land where the predicted variance matches the squared residual.
- How would you spot the variance shortcut in a training run?Track the two loss components separately, not just their sum. The tell is a falling total loss driven almost entirely by the `s` term while the plain squared residual stays flat, together with `s` climbing fastest in the regions that were hardest at initialisation. Compare the mean's squared error against a plain squared-error baseline trained on the same data — if it is much worse, the variance head is absorbing the work.
- Why does the mean fitted this way differ from a plain squared-error fit?Because each example's mean gradient is scaled by one over its predicted variance, the fit is precision-weighted rather than uniform. Low-noise regions dominate the shared trunk, so the model allocates capacity there and tolerates larger residuals in high-noise regions. That is usually the behaviour you want, but it means an evaluation that averages plain squared residuals over all regions will not necessarily favour it.
- Is a two-output head the only way to let the loss express varying noise scale?No. If the noise scale is known or well-approximated from metadata — sensor class, aggregation window, sample count behind a target — you can supply per-example weights directly and keep an ordinary squared-error head, which avoids the extra optimisation failure mode entirely. The learned head earns its place when the scale varies with structure you cannot enumerate in advance.
saying these in an interview costs you the question
- Says the network outputs a variance directly with no positivity constraint
- Omits the log-variance term and does not notice the loss degenerates
- Thinks the predicted variance affects only the variance head's gradient
- Believes the mean gradient is unchanged from plain squared error
- Never mentions that inflating the variance is an easy descent direction