skip to content

How does a Gumbel-Softmax relaxation make a categorical latent trainable by gradients?

level: seniorimportance: should knowfreq 32%

answer

  1. put all randomness in parameter-free noise
  2. Gumbel noise added to the logits
  3. argmax is the non-differentiable part
  4. replace argmax with a tempered softmax
  5. low temperature: sharper sample, noisier gradient

basics

~20 s

Gumbel-Softmax adds independent Gumbel noise to the logits and replaces the argmax with a temperature-scaled softmax. The noise is parameter-free, so gradients flow to the logits; the price is a soft sample and a biased gradient.

solid answer

~50 s

Start from the Gumbel-max trick: adding independent standard Gumbel noise to the log-probabilities and taking the argmax draws an exact categorical sample. The argmax is the non-differentiable part, so replace it with a softmax divided by a temperature. The result is a point on the simplex whose randomness comes entirely from noise that no parameter touches, so it is a reparameterized sample and the logits get a pathwise gradient. Temperature is the knob: high temperature gives a smooth, near-uniform mixture with well-behaved gradients, low temperature gives a near-one-hot draw whose distribution approaches the true categorical but whose gradients get noisy and saturated. Any non-zero temperature makes the sample a soft mixture rather than a genuine discrete draw, so the gradient is biased with respect to the discrete objective. In practice you anneal from soft toward a small non-zero floor.

go deeper

for a junior

Know the shape of the idea: noise is added to the class scores, a softmax with a temperature stands in for the argmax, and that substitution is what lets gradients reach the scores at all.

for a middle

Explain why the argmax is the blocking operation, why the noise must not depend on the logits, and what moving the temperature up or down does to the sample and to the gradient.

for a senior

Show operational judgment: an annealing schedule you would actually run, how you monitor sample entropy, and how you handle the gap between soft training inputs and a hard argmax at inference.

for a principal

Own the decision to accept a biased gradient at all. Weigh a permanently biased objective plus a schedule every future team must maintain against the alternatives, and say what evidence would make you reverse the call.

## Why a categorical latent is hard A reparameterized draw needs the sample to be a differentiable function of the parameters given parameter-free noise. For a categorical latent that is impossible: any map from continuous noise to one of `K` classes is piecewise constant, so its derivative with respect to the class logits is zero almost everywhere and undefined at the boundaries. Nudging a logit does not move the sample at all until it crosses a threshold, at which point it jumps. ## The Gumbel-max trick There is a classical way to draw a categorical sample using only independent per-class noise. If `g_1 ... g_K` are independent standard Gumbel variates and `a_1 ... a_K` are the class logits, then ``` index = argmax_i (a_i + g_i) ``` is distributed exactly as the categorical with probabilities `softmax(a)`. Gumbel noise itself is easy to produce from a uniform draw `u` as `g = -log(-log(u))`. This is already progress of a specific kind: it moves *all* the randomness into noise that does not depend on the logits. Everything else is a deterministic function of the logits. The only obstacle left is the argmax. ## Relaxing the argmax Gumbel-Softmax (also published as the Concrete distribution) replaces the argmax with a temperature-scaled softmax, producing a vector `y` on the simplex: ``` y_i = exp((a_i + g_i)/tau) / sum_j exp((a_j + g_j)/tau) ``` Since the noise is parameter-free and the softmax is smooth, `y` is a reparameterized sample: the logits receive an ordinary pathwise gradient. The decoder consumes `y`, which behaves like a soft, weighted blend of the classes. **What the temperature controls.** As `tau` grows, `y` flattens toward uniform: the sample loses its dependence on the logits, but the gradient is smooth and low-variance. As `tau` falls toward zero, the softmax sharpens and `y` approaches a one-hot vector; the distribution of `y` converges to the true categorical, but the softmax saturates and gradient estimates become high-variance and ill-conditioned. This is a bias-variance dial in the most literal sense, and there is no setting that gives you both ends. **The bias.** At any non-zero temperature, the thing you sampled is a mixture, not a class. The gradient is therefore an unbiased estimate of the gradient of a *relaxed* objective, not of the discrete one you actually care about. That is a real cost, not a formality, and it is what separates this approach from the score-function estimator, which is unbiased on the exact discrete objective but far noisier. ## Annealing, and the train/test gap The standard recipe is a schedule: start warm enough that gradients are informative, anneal toward a small non-zero floor over training, and never go to exactly zero, which is numerically hopeless. The floor is a real hyperparameter and interacts with the learning rate — a schedule that sharpens faster than the model learns will freeze the assignment early. The second operational issue is that at inference you almost always want a true discrete choice, so you take the argmax. If training only ever showed the decoder blends, the decoder has never seen a genuine one-hot input and reconstruction quality can drop noticeably at the switch. The straight-through variant addresses this by discretising the sample in the forward pass — passing the hard one-hot vector to the decoder — while using the soft vector's derivative in the backward pass. Forward-time behaviour then matches inference exactly, at the price of a gradient that no longer corresponds to the function actually computed. Which of the two is better is empirical. ## Scale: a large catalogue Suppose the latent choice is one item out of a 10,000-item catalogue. Two things bite. First, the trick needs one noise variate per class per sample, so you draw and softmax over 10,000 values every step. Second, and worse, a soft sample early in training is a weighted average over 10,000 item embeddings — a vector that corresponds to no item, sitting near the centroid of the catalogue, from which the decoder learns very little. The signal only becomes meaningful as the temperature drops and the mixture concentrates on a handful of items. Practitioners usually respond by annealing more aggressively, restricting the relaxation to a shortlist of plausible classes, or abandoning the relaxation for a scheme whose forward pass is discrete from the start. ## What to remember The relaxation buys you a low-variance pathwise gradient into a categorical choice, and it charges you bias, a temperature schedule to tune, and a mismatch between what training and inference feed the decoder. Those three costs are the whole conversation in an interview.

  • How would you schedule the temperature over a training run?
    Start warm enough that the mixture is genuinely soft and the gradients are informative, then anneal geometrically toward a small non-zero floor; going to exactly zero is numerically hopeless. Treat the floor and the annealing rate as real hyperparameters, and watch the entropy of the sampled vectors: if it collapses long before the model has learned, the schedule is sharpening faster than the encoder can adapt.
  • What does the straight-through variant change?
    It sends a hard one-hot vector to the decoder in the forward pass while using the soft relaxed vector's derivative in the backward pass. Training then matches inference exactly, which removes the blended-input mismatch, but the gradient no longer corresponds to the function actually computed, so it is a differently biased estimator rather than a less biased one.
  • What does a 10,000-class catalogue cost you in this scheme?
    Two things. You draw and normalise over 10,000 noise variates every sample, which is pure overhead. More damaging, a warm soft sample is a weighted average over 10,000 embeddings, a vector near the catalogue centroid that resembles no real item, so early training signal is weak. Aggressive annealing, a candidate shortlist, or a discrete-forward scheme are the usual answers.

saying these in an interview costs you the question

  • Claims the relaxed gradient is unbiased for the discrete objective
  • Thinks higher temperature makes samples more one-hot
  • Says the Gumbel noise depends on the logits
  • Ignores the mismatch between soft training and hard inference
  • Treats the temperature as a fixed constant needing no schedule

context