In uncertainty weighting of a multi-task loss, how are the per-task weights learned and what stops them collapsing to zero?
answer
- one learned scalar per task
- the weight is an exponential of it
- an extra term punishes claiming noise
- stationary point puts the weight at 1/L
basics
~20 sEach task carries a learned scalar noise parameter, and its loss weight is the inverse of that scalar, so hard-to-fit tasks are downweighted automatically. An added log-noise penalty per task is what stops every weight sliding to zero.
solid answer
~50 sRather than hand-tuning weights, you give each task one learned scalar and optimise it jointly with the network. Writing `s_i` for task i's learned log-variance, the objective becomes `L = sum_i ( exp(-s_i) * L_i + s_i )`, up to a factor of a half on regression terms. The exponential is the effective weight: a task the model cannot fit pushes `s_i` up and shrinks its own influence. The bare `+ s_i` is the guard — without it the optimiser drives every `s_i` upward, all weights go to zero and the total goes to zero with nothing learned. Setting the derivative in `s_i` to zero gives `exp(-s_i) = 1 / L_i`, so at the stationary point every term contributes about one and the terms self-normalise by magnitude. Parameterising by log-variance keeps the weight positive and the optimisation well scaled.
code
python · 14 linesimport math
def learned_weights(losses, with_log_term, steps=2000, lr=0.05):
s = [0.0] * len(losses) # s[i] = learned log-variance of task i
for _ in range(steps):
for i, L in enumerate(losses):
# d/ds of exp(-s) * L + s is -exp(-s) * L + 1
grad = -math.exp(-s[i]) * L + (1.0 if with_log_term else 0.0)
s[i] -= lr * grad
return [round(math.exp(-x), 3) for x in s] # effective weights
losses = [4.0, 0.5]
print(learned_weights(losses, with_log_term=True)) # settles at 1/L: [0.25, 2.0]
print(learned_weights(losses, with_log_term=False)) # both collapse toward 0.0go deeper
Know that loss weights need not be constants you choose by hand — they can be extra scalar parameters, one per task, that the optimiser trains alongside the network's weights.
Be able to write the objective with a learned log-variance per task, say that the exponential of its negative is the effective weight, and explain what the extra additive log term is there to prevent.
Show that you would log the learned weights over training, spot a task being progressively abandoned, and articulate that balancing by loss magnitude is not the same as balancing by what each task needs.
Frame it as trading a hyperparameter sweep for handing the objective's trade-off to the optimiser, and decide when a product-critical task must keep a weight the model is not permitted to shrink.
## The idea Hand-tuning the weights of a multi-task sum is a search whose cost grows with the number of tasks. Uncertainty weighting removes the search by making each weight a **parameter of the model**, trained by the same optimiser as everything else. Each task gets one scalar — not per example, one per task for the whole dataset — that plays the role of a noise scale for that task, and the weight is derived from it. The scalar is stored as a log-variance, written `s`, rather than as a variance or a standard deviation. Two reasons: `exp(-s)` is positive for any real `s`, so the optimiser can never produce a negative or zero weight; and the quantity being optimised stays on a well-conditioned scale instead of hugging a hard boundary at zero. ## The objective For tasks with losses `L_1 ... L_k` and learned log-variances `s_1 ... s_k`: ``` L_total = sum_i ( exp(-s_i) * L_i + s_i ) ``` (The usual derivation, which treats each task's loss as a log-likelihood under a learned homoscedastic noise scale, puts a factor of a half on squared-error terms; the structure is identical and the half changes nothing qualitative.) Two pieces: - `exp(-s_i)` is the **effective weight**. Large `s_i` means a large assumed noise scale for that task and a small weight. - `+ s_i` is the **penalty on claiming noise**. It is the entire reason the scheme is not degenerate. ## Why the penalty is load-bearing Suppose you drop the `+ s_i` term and keep only `exp(-s_i) * L_i`. The optimiser has a trivial way to minimise it: send every `s_i` to infinity. Every weight collapses toward zero, the total collapses toward zero, and the network learns nothing, because a loss that is multiplied by nearly zero produces nearly no gradient for any parameter. The `+ s_i` term rises as the weight falls, so shrinking a task's weight is no longer free. The balance point is exact and worth being able to derive on a whiteboard. Differentiating `exp(-s) * L + s` with respect to `s` gives `-exp(-s) * L + 1`. Setting that to zero gives ``` exp(-s) = 1 / L ``` So the learned weight converges to the reciprocal of that task's current loss, and each weighted term contributes about one to the total. This is the honest description of what the method does: **it normalises tasks by loss magnitude, continuously and automatically, as those magnitudes drift during training.** It is not a statement about which task matters. ## What it buys, and what it does not What it buys: the weight sweep disappears; the balance tracks scale changes over training instead of being frozen at the value that suited step zero; and the learned scalars are cheap — a handful of extra parameters. What it does not buy: - **It does not resolve disagreement in gradient direction.** Two tasks whose gradients point against each other at the shared parameters remain in conflict at any weighting; this method only rescales magnitudes. - **It can abandon a task you care about.** A task with a high irreducible error floor keeps a large loss, so it keeps a small weight — permanently. If that task is the product metric, the method has optimised the wrong thing on your behalf. - **It can be gamed by a degenerate solution.** If a task's loss can be driven low by a trivial prediction, its weight grows, and the model is rewarded for spending shared capacity on the easy win. ## Operating it - **Initialise** the log-variances at zero, so every weight starts at one and the first steps behave like an equal-weight run. - **Exclude them from weight decay.** Decay pulls `s_i` toward zero, which pulls every weight toward one — quietly restoring the fixed equal weighting the method was supposed to replace. - **Log the learned weights as curves.** They are the most informative diagnostic the method gives you: a weight decaying monotonically toward zero is a task being abandoned, and you want to see that while the run is still going. - **Consider pinning the primary task.** A common compromise is to fix the weight of the task the product is judged on and let the auxiliaries' weights be learned, so the optimiser can rebalance the support cast but cannot demote the star. ## Neighbouring approaches Balancing by **loss magnitude** is one choice; balancing by **gradient magnitude** is another. GradNorm adapts task weights so that the gradient norms each task produces at a shared layer stay in a target relationship with the tasks' relative training progress. It attacks the same imbalance one level closer to the update, at the cost of measuring per-task gradients. Both are heuristics for scale, and neither is a substitute for deciding, as a product question, which task is allowed to lose.
- What does a learned weight settle to if you let it converge, and why does that matter?The derivative of `exp(-s) * L + s` vanishes when `exp(-s) = 1 / L`, so the weight settles at the reciprocal of that task's current loss and each term contributes about one. It matters because the method is balancing by loss magnitude alone: a task with a high irreducible error floor keeps a large loss and therefore a permanently small weight, however important it is.
- Should the learned log-variance parameters be subject to weight decay?No. Decay pulls each log-variance toward zero, and a log-variance of zero means a weight of one, so you are quietly dragging the model back to the fixed equal weighting you adopted the method to escape. Exclude those scalars from decay, initialise them at zero, and monitor them as curves so you can see a task being downweighted out of existence.
- When would you not use learned task weights?When one task carries the product metric and the rest are supporting objectives. Loss-magnitude balancing will downweight an intrinsically noisy task precisely because it is noisy, which is the opposite of what you want if that task is what you ship. The usual compromise is to pin the primary task's weight and learn only the auxiliaries', or to hand-tune when there are just two terms.
saying these in an interview costs you the question
- Omits the log term and cannot say why the weights collapse
- Claims learned weights need no monitoring during training
- Says the method also fixes conflicting gradient directions
- Applies weight decay to the learned log-variance parameters
- Assumes a downweighted task must be an unimportant one