Practice / Gradients of expectations

The ELBO and the reparameterisation trick

Ten problems on the evidence lower bound: its derivation from Jensen's inequality, the gap as a KL divergence to the posterior, the reconstruction-minus-KL form, the KL between diagonal Gaussians, the reparameterised gradients in μ, σ and log σ², the variance of the score-function estimator next to the reparameterised one, a linear Gaussian model solved in closed form with its tight bound, the one-sample gradient of a VAE step, and the Monte Carlo KL estimate and its variance, with worked solutions and the mistakes that reverse a KL, detach the sample or apply Jensen the wrong way.

Before you start

A variational autoencoder is trained by maximising a lower bound on the log-likelihood of the data, and the bound exists because of one inequality, Jensen's, applied to one expectation, over a distribution that the model chooses. Everything else is bookkeeping: the bound splits into a reconstruction term and a KL divergence, the KL between Gaussians has a closed form, and the gradient of the reconstruction term with respect to the distribution's own parameters is computed by writing the sample as a deterministic function of noise. These ten problems derive the evidence lower bound, identify the gap exactly, split it both ways, compute the Gaussian KL in the parameterisation encoders use, derive the reparameterised gradient and compare its variance with the score-function alternative, solve a linear Gaussian model where every quantity is closed-form and the bound is tight, and work through the one-sample gradient of a VAE step and the one-sample KL estimate with its variance. The five mistakes at the end are the ones that train but train the wrong thing: the KL written in the other order, a sample treated as a constant, a Monte Carlo KL averaged under the prior, Jensen's inequality reversed, and a sample scaled by the variance instead of the standard deviation.

  • A latent-variable model has a prior p(z)p(z) over a latent zz, a likelihood p(x∣z)p(x \mid z) for the observed xx (in a VAE a decoder with parameters θ\theta, written pθ(x∣z)p_\theta(x \mid z)), the joint p(x,z)=p(z) p(x∣z)p(x, z) = p(z)\,p(x \mid z), the evidence p(x)=∫p(x,z) dzp(x) = \int p(x, z)\,dz (a sum when zz is discrete) and the posterior p(z∣x)=p(x,z)/p(x)p(z \mid x) = p(x, z)/p(x).
  • A variational distribution q(z)q(z) is any distribution over zz that is positive wherever p(x,z)p(x, z) is; in a VAE it is the encoder's output qϕ(z∣x)q_\phi(z \mid x), and for a fixed xx this page writes it q(z)q(z). Eq\mathbb{E}_q is expectation under qq. The evidence lower bound is L(q)=Eq[log⁡p(x,z)−log⁡q(z)]\mathcal{L}(q) = \mathbb{E}_q[\log p(x, z) - \log q(z)].
  • Jensen's inequality for the concave logarithm: log⁡E[w]≥E[log⁡w]\log\mathbb{E}[w] \ge \mathbb{E}[\log w] for a positive random variable ww, with equality exactly when ww is constant.
  • From the entropy page: KL⁡(q ∥ p)=Eq[log⁡q(z)−log⁡p(z)]≥0\operatorname{KL}(q\,\|\,p) = \mathbb{E}_q[\log q(z) - \log p(z)] \ge 0, with equality only when q=pq = p; the entropy is H(q)=−Eq[log⁡q(z)]H(q) = -\mathbb{E}_q[\log q(z)]; and for one-dimensional Gaussians, KL⁡(N(μ1,σ12) ∥ N(μ2,σ22))=log⁡σ2σ1+σ12+(μ1−μ2)22σ22−12\operatorname{KL}\big(\mathcal{N}(\mu_1, \sigma_1^2)\,\|\,\mathcal{N}(\mu_2, \sigma_2^2)\big) = \log\frac{\sigma_2}{\sigma_1} + \frac{\sigma_1^2 + (\mu_1 - \mu_2)^2}{2\sigma_2^2} - \frac12.
  • A diagonal Gaussian N(μ,diag⁡(σ2))\mathcal{N}(\mu, \operatorname{diag}(\sigma^2)) on Rd\mathbb{R}^d has independent coordinates zj∼N(μj,σj2)z_j \sim \mathcal{N}(\mu_j, \sigma_j^2). Encoders output μ\mu and ss with sj=log⁡σj2s_j = \log\sigma_j^2, so σj=esj/2\sigma_j = e^{s_j/2}.
  • The reparameterisation: z=μ+σ⊙εz = \mu + \sigma\odot\varepsilon with ε∼N(0,I)\varepsilon \sim \mathcal{N}(0, I) has the distribution N(μ,diag⁡(σ2))\mathcal{N}(\mu, \operatorname{diag}(\sigma^2)) (the multivariate-Gaussian page's Problem 8 with L=diag⁡(σ)L = \operatorname{diag}(\sigma)). For ε∼N(0,1)\varepsilon \sim \mathcal{N}(0, 1) the moments used below are E[ε]=E[ε3]=E[ε5]=0\mathbb{E}[\varepsilon] = \mathbb{E}[\varepsilon^3] = \mathbb{E}[\varepsilon^5] = 0, E[ε2]=1\mathbb{E}[\varepsilon^2] = 1, E[ε4]=3\mathbb{E}[\varepsilon^4] = 3 and E[ε6]=15\mathbb{E}[\varepsilon^6] = 15.
  • Variances and covariances follow the variance page: Var⁡(u)=E[u2]−(E[u])2\operatorname{Var}(u) = \mathbb{E}[u^2] - (\mathbb{E}[u])^2 and Var⁡(u+v)=Var⁡(u)+Var⁡(v)+2Cov⁡(u,v)\operatorname{Var}(u + v) = \operatorname{Var}(u) + \operatorname{Var}(v) + 2\operatorname{Cov}(u, v).

Builds on: Entropy, cross-entropy and KL divergence, The multivariate Gaussian: gradients and identities

Problems

  1. ·

    Show that log⁡p(x)≥L(q)=Eq[log⁡p(x,z)−log⁡q(z)]\log p(x) \ge \mathcal{L}(q) = \mathbb{E}_q[\log p(x, z) - \log q(z)] for every admissible qq, and say when the two are equal.

  2. ·

    Show that log⁡p(x)−L(q)=KL⁡(q(z) ∥ p(z∣x))\log p(x) - \mathcal{L}(q) = \operatorname{KL}\big(q(z)\,\|\,p(z \mid x)\big).

  3. ··

    Show that L(q)=Eq[log⁡p(x∣z)]−KL⁡(q(z) ∥ p(z))=Eq[log⁡p(x∣z)]+Eq[log⁡p(z)]+H(q)\mathcal{L}(q) = \mathbb{E}_q[\log p(x \mid z)] - \operatorname{KL}\big(q(z)\,\|\,p(z)\big) = \mathbb{E}_q[\log p(x \mid z)] + \mathbb{E}_q[\log p(z)] + H(q).

  4. ··

    Let q=N(μ1,diag⁡(σ12))q = \mathcal{N}(\mu_1, \operatorname{diag}(\sigma_1^2)) and p=N(μ2,diag⁡(σ22))p = \mathcal{N}(\mu_2, \operatorname{diag}(\sigma_2^2)) on Rd\mathbb{R}^d. Show that

    KL⁡(q ∥ p)=∑j=1d(log⁡σ2jσ1j+σ1j2+(μ1j−μ2j)22σ2j2−12),\operatorname{KL}(q\,\|\,p) = \sum_{j=1}^{d}\Big(\log\frac{\sigma_{2j}}{\sigma_{1j}} + \frac{\sigma_{1j}^2 + (\mu_{1j} - \mu_{2j})^2}{2\sigma_{2j}^2} - \frac12\Big),

    and that with μ2=0\mu_2 = 0, σ2=1\sigma_2 = \mathbf{1} and sj=log⁡σ1j2s_j = \log\sigma_{1j}^2 it becomes 12∑j(μ1j2+esj−sj−1)\tfrac12\sum_j\big(\mu_{1j}^2 + e^{s_j} - s_j - 1\big).

  5. ··

    Let z=μ+σ⊙εz = \mu + \sigma\odot\varepsilon with ε∼N(0,I)\varepsilon \sim \mathcal{N}(0, I), and F(μ,σ)=E[f(z)]F(\mu, \sigma) = \mathbb{E}[f(z)] for a smooth f:Rd→Rf: \mathbb{R}^d \to \mathbb{R}. Show that ∇μF=E[∇f(z)]\nabla_\mu F = \mathbb{E}[\nabla f(z)], that ∂F/∂σj=E[∂jf(z) εj]\partial F/\partial\sigma_j = \mathbb{E}[\partial_j f(z)\,\varepsilon_j], and that with sj=log⁡σj2s_j = \log\sigma_j^2, ∂F/∂sj=σj2 E[∂jf(z) εj]\partial F/\partial s_j = \tfrac{\sigma_j}{2}\,\mathbb{E}[\partial_j f(z)\,\varepsilon_j]. Verify all three on f(z)=z⊤Bz+c⊤zf(z) = z^\top Bz + c^\top z with BB symmetric, where F=μ⊤Bμ+∑jBjjσj2+c⊤μF = \mu^\top B\mu + \sum_j B_{jj}\sigma_j^2 + c^\top\mu.

  6. ···

    Take q=N(μ,σ2)q = \mathcal{N}(\mu, \sigma^2) in one dimension and f(z)=z2f(z) = z^2. Two unbiased estimators of ∂ Eq[f(z)]/∂μ\partial\,\mathbb{E}_q[f(z)]/\partial\mu are the reparameterised gR=f′(μ+σε)=2(μ+σε)g_R = f'(\mu + \sigma\varepsilon) = 2(\mu + \sigma\varepsilon) and the score-function gS=f(z) ∂μlog⁡q(z)=z2(z−μ)/σ2g_S = f(z)\,\partial_\mu\log q(z) = z^2(z - \mu)/\sigma^2 with z=μ+σεz = \mu + \sigma\varepsilon. Show both have mean 2μ2\mu and compute their variances.

  7. ··

    A linear Gaussian model: p(z)=N(0,1)p(z) = \mathcal{N}(0, 1), p(x∣z)=N(wz+b,γ2)p(x \mid z) = \mathcal{N}(wz + b, \gamma^2) with w,b,γw, b, \gamma fixed, and q(z)=N(μ,σ2)q(z) = \mathcal{N}(\mu, \sigma^2). Show that

    L(μ,σ2)=−12log⁡(2πγ2)−(x−wμ−b)2+w2σ22γ2−12(μ2+σ2−log⁡σ2−1).\mathcal{L}(\mu, \sigma^2) = -\tfrac12\log(2\pi\gamma^2) - \frac{(x - w\mu - b)^2 + w^2\sigma^2}{2\gamma^2} - \tfrac12\big(\mu^2 + \sigma^2 - \log\sigma^2 - 1\big).
  8. ···

    Maximise Problem 7's L\mathcal{L} over μ\mu and σ2\sigma^2. Show that μ∗=w(x−b)γ2+w2\mu^* = \dfrac{w(x - b)}{\gamma^2 + w^2} and σ∗2=γ2γ2+w2\sigma^{*2} = \dfrac{\gamma^2}{\gamma^2 + w^2}, that q∗=N(μ∗,σ∗2)q^* = \mathcal{N}(\mu^*, \sigma^{*2}) is the exact posterior p(z∣x)p(z \mid x), and that the bound is tight: L∗=log⁡p(x)\mathcal{L}^* = \log p(x) with p(x)=N(x;b,w2+γ2)p(x) = \mathcal{N}(x; b, w^2 + \gamma^2).

  9. ···

    Let the decoder give log⁡pθ(x∣z)=g(z)\log p_\theta(x \mid z) = g(z) for a differentiable gg, with q=N(μ,σ2)q = \mathcal{N}(\mu, \sigma^2) and prior N(0,1)\mathcal{N}(0, 1). The one-sample estimate of the ELBO is L^(μ,σ)=g(μ+σε)−12(μ2+σ2−log⁡σ2−1)\hat{\mathcal{L}}(\mu, \sigma) = g(\mu + \sigma\varepsilon) - \tfrac12\big(\mu^2 + \sigma^2 - \log\sigma^2 - 1\big) with ε\varepsilon drawn once. Compute ∂L^/∂μ\partial\hat{\mathcal{L}}/\partial\mu, ∂L^/∂σ\partial\hat{\mathcal{L}}/\partial\sigma and ∂L^/∂s\partial\hat{\mathcal{L}}/\partial s with s=log⁡σ2s = \log\sigma^2, and show that each has expectation equal to the corresponding gradient of L\mathcal{L}.

  10. ···

    Instead of Problem 4's closed form, the KL can be estimated from one sample: K^=log⁡q(z)−log⁡p(z)\hat K = \log q(z) - \log p(z) with z=μ+σεz = \mu + \sigma\varepsilon, q=N(μ,σ2)q = \mathcal{N}(\mu, \sigma^2) and p=N(0,1)p = \mathcal{N}(0, 1). Show that K^=−log⁡σ+12μ2+μσε+12(σ2−1)ε2\hat K = -\log\sigma + \tfrac12\mu^2 + \mu\sigma\varepsilon + \tfrac12(\sigma^2 - 1)\varepsilon^2, that E[K^]=KL⁡(q ∥ p)\mathbb{E}[\hat K] = \operatorname{KL}(q\,\|\,p), and that Var⁡(K^)=μ2σ2+12(σ2−1)2\operatorname{Var}(\hat K) = \mu^2\sigma^2 + \tfrac12(\sigma^2 - 1)^2.

Worked solutions

Problem 1

Show that log⁡p(x)≥L(q)=Eq[log⁡p(x,z)−log⁡q(z)]\log p(x) \ge \mathcal{L}(q) = \mathbb{E}_q[\log p(x, z) - \log q(z)] for every admissible qq, and say when the two are equal.

  1. p(x)=∫p(x,z) dz=∫q(z) p(x,z)q(z) dz=Eq[p(x,z)q(z)]p(x) = \int p(x, z)\,dz = \int q(z)\,\dfrac{p(x, z)}{q(z)}\,dz = \mathbb{E}_q\Big[\dfrac{p(x, z)}{q(z)}\Big].Multiply and divide the integrand by q(z)q(z), which is positive wherever p(x,z)p(x, z) is; an integral against qq is an expectation under qq.
  2. log⁡p(x)=log⁡Eq[w]≥Eq[log⁡w]\log p(x) = \log\mathbb{E}_q[w] \ge \mathbb{E}_q[\log w] with w=p(x,z)/q(z)w = p(x, z)/q(z).Jensen's inequality for the concave logarithm.
  3. Eq[log⁡w]=Eq[log⁡p(x,z)−log⁡q(z)]=L(q)\mathbb{E}_q[\log w] = \mathbb{E}_q[\log p(x, z) - \log q(z)] = \mathcal{L}(q).The log of a quotient is a difference.
  4. Equality holds exactly when ww is constant under qq: p(x,z)=c q(z)p(x, z) = c\,q(z), and integrating over zz gives c=p(x)c = p(x), so q(z)=p(x,z)/p(x)=p(z∣x)q(z) = p(x, z)/p(x) = p(z \mid x).Jensen's equality condition for a strictly concave function; qq integrates to 11.
  5. log⁡p(x)≥L(q)=Eq[log⁡p(x,z)−log⁡q(z)]\log p(x) \ge \mathcal{L}(q) = \mathbb{E}_q[\log p(x, z) - \log q(z)], with equality exactly when q(z)=p(z∣x)q(z) = p(z \mid x)The bound holds for every qq, which is what makes it usable: pick a family that can be sampled and differentiated, and maximise over it. Problem 2 computes the gap exactly, and Mistake 4 is this derivation with Jensen's inequality the wrong way round.

Problem 2

Show that log⁡p(x)−L(q)=KL⁡(q(z) ∥ p(z∣x))\log p(x) - \mathcal{L}(q) = \operatorname{KL}\big(q(z)\,\|\,p(z \mid x)\big).

  1. log⁡p(x,z)=log⁡p(z∣x)+log⁡p(x)\log p(x, z) = \log p(z \mid x) + \log p(x).p(x,z)=p(z∣x) p(x)p(x, z) = p(z \mid x)\,p(x), the product rule of the Bayes page.
  2. L(q)=Eq[log⁡p(z∣x)+log⁡p(x)−log⁡q(z)]=log⁡p(x)+Eq[log⁡p(z∣x)−log⁡q(z)]\mathcal{L}(q) = \mathbb{E}_q[\log p(z \mid x) + \log p(x) - \log q(z)] = \log p(x) + \mathbb{E}_q[\log p(z \mid x) - \log q(z)].Substitute; log⁡p(x)\log p(x) does not depend on zz, and the expectation of a constant is the constant.
  3. Eq[log⁡p(z∣x)−log⁡q(z)]=−Eq[log⁡q(z)−log⁡p(z∣x)]=−KL⁡(q ∥ p(z∣x))\mathbb{E}_q[\log p(z \mid x) - \log q(z)] = -\mathbb{E}_q[\log q(z) - \log p(z \mid x)] = -\operatorname{KL}\big(q\,\|\,p(z \mid x)\big).The definition of KL with the posterior as the second argument.
  4. log⁡p(x)=L(q)+KL⁡(q(z) ∥ p(z∣x))\log p(x) = \mathcal{L}(q) + \operatorname{KL}\big(q(z)\,\|\,p(z \mid x)\big), so log⁡p(x)−L(q)=KL⁡(q ∥ p(z∣x))≥0\log p(x) - \mathcal{L}(q) = \operatorname{KL}\big(q\,\|\,p(z \mid x)\big) \ge 0KL⁡≥0\operatorname{KL} \ge 0 (the entropy page's Problem 5) recovers Problem 1 with the gap named. Since log⁡p(x)\log p(x) does not depend on qq, raising L\mathcal{L} over qq is the same as lowering the KL to the true posterior, without ever computing that posterior or its normaliser p(x)p(x). Problem 8 shows the gap closing to 00 when the family contains the posterior.

Problem 3

Show that L(q)=Eq[log⁡p(x∣z)]−KL⁡(q(z) ∥ p(z))=Eq[log⁡p(x∣z)]+Eq[log⁡p(z)]+H(q)\mathcal{L}(q) = \mathbb{E}_q[\log p(x \mid z)] - \operatorname{KL}\big(q(z)\,\|\,p(z)\big) = \mathbb{E}_q[\log p(x \mid z)] + \mathbb{E}_q[\log p(z)] + H(q).

  1. log⁡p(x,z)=log⁡p(x∣z)+log⁡p(z)\log p(x, z) = \log p(x \mid z) + \log p(z).The product rule the other way round.
  2. L(q)=Eq[log⁡p(x∣z)]+Eq[log⁡p(z)]−Eq[log⁡q(z)]\mathcal{L}(q) = \mathbb{E}_q[\log p(x \mid z)] + \mathbb{E}_q[\log p(z)] - \mathbb{E}_q[\log q(z)].Substitute into the definition and use linearity of Eq\mathbb{E}_q.
  3. Eq[log⁡p(z)]−Eq[log⁡q(z)]=−Eq[log⁡q(z)−log⁡p(z)]=−KL⁡(q ∥ p(z))\mathbb{E}_q[\log p(z)] - \mathbb{E}_q[\log q(z)] = -\mathbb{E}_q[\log q(z) - \log p(z)] = -\operatorname{KL}\big(q\,\|\,p(z)\big).The definition of KL with the prior as the second argument.
  4. −Eq[log⁡q(z)]=H(q)-\mathbb{E}_q[\log q(z)] = H(q).The definition of entropy.
  5. L(q)=Eq[log⁡p(x∣z)]−KL⁡(q(z) ∥ p(z))=Eq[log⁡p(x∣z)]+Eq[log⁡p(z)]+H(q)\mathcal{L}(q) = \mathbb{E}_q[\log p(x \mid z)] - \operatorname{KL}\big(q(z)\,\|\,p(z)\big) = \mathbb{E}_q[\log p(x \mid z)] + \mathbb{E}_q[\log p(z)] + H(q)The first form is the VAE objective: a reconstruction term, the expected log-likelihood of xx under the decoder, minus a regulariser that pulls qq towards the prior. For Gaussian qq and prior the KL is closed-form (Problem 4), so only the reconstruction term needs sampling (Problem 9). The second form shows what the entropy does: without H(q)H(q), maximising Eq[log⁡p(x,z)]\mathbb{E}_q[\log p(x, z)] over qq would collapse qq onto the single zz that maximises p(x,z)p(x, z), and the bound would no longer be a bound.

Problem 4

Let q=N(μ1,diag⁡(σ12))q = \mathcal{N}(\mu_1, \operatorname{diag}(\sigma_1^2)) and p=N(μ2,diag⁡(σ22))p = \mathcal{N}(\mu_2, \operatorname{diag}(\sigma_2^2)) on Rd\mathbb{R}^d. Show that

KL⁡(q ∥ p)=∑j=1d(log⁡σ2jσ1j+σ1j2+(μ1j−μ2j)22σ2j2−12),\operatorname{KL}(q\,\|\,p) = \sum_{j=1}^{d}\Big(\log\frac{\sigma_{2j}}{\sigma_{1j}} + \frac{\sigma_{1j}^2 + (\mu_{1j} - \mu_{2j})^2}{2\sigma_{2j}^2} - \frac12\Big),

and that with μ2=0\mu_2 = 0, σ2=1\sigma_2 = \mathbf{1} and sj=log⁡σ1j2s_j = \log\sigma_{1j}^2 it becomes 12∑j(μ1j2+esj−sj−1)\tfrac12\sum_j\big(\mu_{1j}^2 + e^{s_j} - s_j - 1\big).

  1. q(z)=∏jqj(zj)q(z) = \prod_j q_j(z_j) with qj=N(μ1j,σ1j2)q_j = \mathcal{N}(\mu_{1j}, \sigma_{1j}^2), and likewise p(z)=∏jpj(zj)p(z) = \prod_j p_j(z_j).A diagonal Gaussian has independent coordinates, so its density is the product of the one-dimensional densities.
  2. log⁡q(z)−log⁡p(z)=∑j(log⁡qj(zj)−log⁡pj(zj))\log q(z) - \log p(z) = \sum_j\big(\log q_j(z_j) - \log p_j(z_j)\big).The log of a product is a sum.
  3. Eq[log⁡qj(zj)−log⁡pj(zj)]=KL⁡(qj ∥ pj)\mathbb{E}_q\big[\log q_j(z_j) - \log p_j(z_j)\big] = \operatorname{KL}(q_j\,\|\,p_j).The expectation of a function of zjz_j alone uses only the marginal of zjz_j under qq, which is qjq_j.
  4. KL⁡(q ∥ p)=∑jKL⁡(qj ∥ pj)=∑j(log⁡σ2jσ1j+σ1j2+(μ1j−μ2j)22σ2j2−12)\operatorname{KL}(q\,\|\,p) = \sum_j\operatorname{KL}(q_j\,\|\,p_j) = \sum_j\Big(\log\dfrac{\sigma_{2j}}{\sigma_{1j}} + \dfrac{\sigma_{1j}^2 + (\mu_{1j} - \mu_{2j})^2}{2\sigma_{2j}^2} - \dfrac12\Big).Linearity of Eq\mathbb{E}_q over step 2, then the entropy page's one-dimensional formula for each pair.
  5. With μ2j=0\mu_{2j} = 0, σ2j=1\sigma_{2j} = 1 and σ1j2=esj\sigma_{1j}^2 = e^{s_j}: log⁡1σ1j=−12sj\log\dfrac{1}{\sigma_{1j}} = -\tfrac12 s_j and σ1j2+μ1j22−12=12(esj+μ1j2−1)\dfrac{\sigma_{1j}^2 + \mu_{1j}^2}{2} - \dfrac12 = \tfrac12\big(e^{s_j} + \mu_{1j}^2 - 1\big).log⁡σ1j=12log⁡σ1j2=12sj\log\sigma_{1j} = \tfrac12\log\sigma_{1j}^2 = \tfrac12 s_j.
  6. KL⁡(q ∥ p)=∑j(log⁡σ2jσ1j+σ1j2+(μ1j−μ2j)22σ2j2−12)\operatorname{KL}(q\,\|\,p) = \sum_j\Big(\log\dfrac{\sigma_{2j}}{\sigma_{1j}} + \dfrac{\sigma_{1j}^2 + (\mu_{1j} - \mu_{2j})^2}{2\sigma_{2j}^2} - \dfrac12\Big); against N(0,I)\mathcal{N}(0, I) with sj=log⁡σ1j2s_j = \log\sigma_{1j}^2 it is 12∑j(μ1j2+esj−sj−1)\tfrac12\sum_j\big(\mu_{1j}^2 + e^{s_j} - s_j - 1\big)The KL between two product distributions is the sum of the coordinate KLs, for any product distributions, not only Gaussians. The check computes each coordinate's expectation by quadrature and also the full multivariate formula 12[tr⁡(Σ2−1Σ1)+(μ2−μ1)⊤Σ2−1(μ2−μ1)−d+log⁡det⁡Σ2−log⁡det⁡Σ1]\tfrac12\big[\operatorname{tr}(\Sigma_2^{-1}\Sigma_1) + (\mu_2 - \mu_1)^\top\Sigma_2^{-1}(\mu_2 - \mu_1) - d + \log\det\Sigma_2 - \log\det\Sigma_1\big] with diagonal Σ\Sigma's; all three agree. The log⁡σ2\log\sigma^2 parameterisation keeps σj>0\sigma_j > 0 without a constraint; the gradient in ss is on the entropy page (Problem 9) and reappears in Problem 9 here.

Problem 5

Let z=μ+σ⊙εz = \mu + \sigma\odot\varepsilon with ε∼N(0,I)\varepsilon \sim \mathcal{N}(0, I), and F(μ,σ)=E[f(z)]F(\mu, \sigma) = \mathbb{E}[f(z)] for a smooth f:Rd→Rf: \mathbb{R}^d \to \mathbb{R}. Show that ∇μF=E[∇f(z)]\nabla_\mu F = \mathbb{E}[\nabla f(z)], that ∂F/∂σj=E[∂jf(z) εj]\partial F/\partial\sigma_j = \mathbb{E}[\partial_j f(z)\,\varepsilon_j], and that with sj=log⁡σj2s_j = \log\sigma_j^2, ∂F/∂sj=σj2 E[∂jf(z) εj]\partial F/\partial s_j = \tfrac{\sigma_j}{2}\,\mathbb{E}[\partial_j f(z)\,\varepsilon_j]. Verify all three on f(z)=z⊤Bz+c⊤zf(z) = z^\top Bz + c^\top z with BB symmetric, where F=μ⊤Bμ+∑jBjjσj2+c⊤μF = \mu^\top B\mu + \sum_j B_{jj}\sigma_j^2 + c^\top\mu.

  1. F(μ,σ)=Eε[f(μ+σ⊙ε)]F(\mu, \sigma) = \mathbb{E}_\varepsilon\big[f(\mu + \sigma\odot\varepsilon)\big], where the distribution of ε\varepsilon involves neither μ\mu nor σ\sigma.The reparameterisation: the parameters have moved from the measure into the integrand, which is what allows the next step.
  2. ∇μF=Eε[∇μf(μ+σ⊙ε)]=E[∇f(z)]\nabla_\mu F = \mathbb{E}_\varepsilon\big[\nabla_\mu f(\mu + \sigma\odot\varepsilon)\big] = \mathbb{E}[\nabla f(z)].Differentiate under the expectation, which is legitimate for a smooth ff whose gradient has finite expectation; ∂z/∂μ=I\partial z/\partial\mu = I.
  3. ∂F∂σj=E[∂jf(z) ∂zj∂σj]=E[∂jf(z) εj]\dfrac{\partial F}{\partial\sigma_j} = \mathbb{E}\Big[\partial_j f(z)\,\dfrac{\partial z_j}{\partial\sigma_j}\Big] = \mathbb{E}[\partial_j f(z)\,\varepsilon_j].Only zjz_j depends on σj\sigma_j, with ∂zj/∂σj=εj\partial z_j/\partial\sigma_j = \varepsilon_j; chain rule.
  4. σj=esj/2\sigma_j = e^{s_j/2} gives ∂σj/∂sj=12esj/2=σj2\partial\sigma_j/\partial s_j = \tfrac12 e^{s_j/2} = \tfrac{\sigma_j}{2}, so ∂F∂sj=σj2 E[∂jf(z) εj]\dfrac{\partial F}{\partial s_j} = \dfrac{\sigma_j}{2}\,\mathbb{E}[\partial_j f(z)\,\varepsilon_j].Chain rule through σj(sj)\sigma_j(s_j).
  5. For the quadratic, E[zz⊤]=μμ⊤+diag⁡(σ2)\mathbb{E}[zz^\top] = \mu\mu^\top + \operatorname{diag}(\sigma^2), so F=tr⁡(B E[zz⊤])+c⊤μ=μ⊤Bμ+∑jBjjσj2+c⊤μF = \operatorname{tr}\big(B\,\mathbb{E}[zz^\top]\big) + c^\top\mu = \mu^\top B\mu + \sum_j B_{jj}\sigma_j^2 + c^\top\mu.z⊤Bz=tr⁡(Bzz⊤)z^\top Bz = \operatorname{tr}(Bzz^\top) and E\mathbb{E} is linear; the variance page's E[zz⊤]=Σ+μμ⊤\mathbb{E}[zz^\top] = \Sigma + \mu\mu^\top with Σ=diag⁡(σ2)\Sigma = \operatorname{diag}(\sigma^2); tr⁡(Bdiag⁡(σ2))=∑jBjjσj2\operatorname{tr}(B\operatorname{diag}(\sigma^2)) = \sum_j B_{jj}\sigma_j^2.
  6. Differentiating FF directly: ∇μF=2Bμ+c\nabla_\mu F = 2B\mu + c, ∂F/∂σj=2Bjjσj\partial F/\partial\sigma_j = 2B_{jj}\sigma_j and ∂F/∂sj=Bjjσj2\partial F/\partial s_j = B_{jj}\sigma_j^2. The estimators: ∇f(z)=2Bz+c\nabla f(z) = 2Bz + c has expectation 2Bμ+c2B\mu + c, and ∂jf(z) εj=(2∑kBjk(μk+σkεk)+cj)εj\partial_j f(z)\,\varepsilon_j = \big(2\sum_k B_{jk}(\mu_k + \sigma_k\varepsilon_k) + c_j\big)\varepsilon_j has expectation 2Bjjσj2B_{jj}\sigma_j.The matrix-calculus page for ∇μ(μ⊤Bμ)=2Bμ\nabla_\mu(\mu^\top B\mu) = 2B\mu; E[εj]=0\mathbb{E}[\varepsilon_j] = 0 kills the μ\mu and cc terms, and E[εkεj]=1\mathbb{E}[\varepsilon_k\varepsilon_j] = 1 for k=jk = j and 00 otherwise keeps one term of the sum.
  7. ∇μF=E[∇f(z)]\nabla_\mu F = \mathbb{E}[\nabla f(z)], ∂F∂σj=E[∂jf(z) εj]\dfrac{\partial F}{\partial\sigma_j} = \mathbb{E}[\partial_j f(z)\,\varepsilon_j], ∂F∂sj=σj2 E[∂jf(z) εj]\dfrac{\partial F}{\partial s_j} = \dfrac{\sigma_j}{2}\,\mathbb{E}[\partial_j f(z)\,\varepsilon_j]; for the quadratic both routes give 2Bμ+c2B\mu + c, 2Bjjσj2B_{jj}\sigma_j and Bjjσj2B_{jj}\sigma_j^2This is the reparameterisation gradient: draw ε\varepsilon, form zz, backpropagate ∇f(z)\nabla f(z) through z=μ+σ⊙εz = \mu + \sigma\odot\varepsilon, so that μ\mu receives ∇f(z)\nabla f(z) and σ\sigma receives ∇f(z)⊙ε\nabla f(z)\odot\varepsilon. One draw gives an unbiased estimate of each gradient; the check averages 400,000400{,}000 draws and matches to Monte Carlo accuracy. In a VAE, ff is the decoder's log⁡pθ(x∣z)\log p_\theta(x \mid z), and the alternative that differentiates the density instead (Problem 6) is far noisier.

Problem 6

Take q=N(μ,σ2)q = \mathcal{N}(\mu, \sigma^2) in one dimension and f(z)=z2f(z) = z^2. Two unbiased estimators of ∂ Eq[f(z)]/∂μ\partial\,\mathbb{E}_q[f(z)]/\partial\mu are the reparameterised gR=f′(μ+σε)=2(μ+σε)g_R = f'(\mu + \sigma\varepsilon) = 2(\mu + \sigma\varepsilon) and the score-function gS=f(z) ∂μlog⁡q(z)=z2(z−μ)/σ2g_S = f(z)\,\partial_\mu\log q(z) = z^2(z - \mu)/\sigma^2 with z=μ+σεz = \mu + \sigma\varepsilon. Show both have mean 2μ2\mu and compute their variances.

  1. ∂∂μEq[z2]=∂∂μ(μ2+σ2)=2μ\dfrac{\partial}{\partial\mu}\mathbb{E}_q[z^2] = \dfrac{\partial}{\partial\mu}(\mu^2 + \sigma^2) = 2\mu.E[z2]=Var⁡(z)+(E[z])2\mathbb{E}[z^2] = \operatorname{Var}(z) + (\mathbb{E}[z])^2, the variance page's Problem 1.
  2. E[gR]=2μ+2σ E[ε]=2μ\mathbb{E}[g_R] = 2\mu + 2\sigma\,\mathbb{E}[\varepsilon] = 2\mu and Var⁡(gR)=4σ2Var⁡(ε)=4σ2\operatorname{Var}(g_R) = 4\sigma^2\operatorname{Var}(\varepsilon) = 4\sigma^2.E[ε]=0\mathbb{E}[\varepsilon] = 0 and Var⁡(ε)=1\operatorname{Var}(\varepsilon) = 1; the variance page's Problem 1 for the scaling.
  3. gS=(μ+σε)2 εσ=μ2ε+2μσε2+σ2ε3σg_S = \dfrac{(\mu + \sigma\varepsilon)^2\,\varepsilon}{\sigma} = \dfrac{\mu^2\varepsilon + 2\mu\sigma\varepsilon^2 + \sigma^2\varepsilon^3}{\sigma}.(z−μ)/σ2=ε/σ(z - \mu)/\sigma^2 = \varepsilon/\sigma; expand the square.
  4. E[gS]=μ2⋅0+2μσ⋅1+σ2⋅0σ=2μ\mathbb{E}[g_S] = \dfrac{\mu^2\cdot 0 + 2\mu\sigma\cdot 1 + \sigma^2\cdot 0}{\sigma} = 2\mu.E[ε]=E[ε3]=0\mathbb{E}[\varepsilon] = \mathbb{E}[\varepsilon^3] = 0 and E[ε2]=1\mathbb{E}[\varepsilon^2] = 1.
  5. E[gS2]=E[(μ2ε+2μσε2+σ2ε3)2]σ2=μ4+12μ2σ2+15σ4+6μ2σ2σ2=μ4σ2+18μ2+15σ2\mathbb{E}[g_S^2] = \dfrac{\mathbb{E}[(\mu^2\varepsilon + 2\mu\sigma\varepsilon^2 + \sigma^2\varepsilon^3)^2]}{\sigma^2} = \dfrac{\mu^4 + 12\mu^2\sigma^2 + 15\sigma^4 + 6\mu^2\sigma^2}{\sigma^2} = \dfrac{\mu^4}{\sigma^2} + 18\mu^2 + 15\sigma^2.Square the three-term sum: the squares give μ4E[ε2]\mu^4\mathbb{E}[\varepsilon^2], 4μ2σ2E[ε4]4\mu^2\sigma^2\mathbb{E}[\varepsilon^4] and σ4E[ε6]\sigma^4\mathbb{E}[\varepsilon^6]; of the cross terms only 2μ2σ2E[ε4]2\mu^2\sigma^2\mathbb{E}[\varepsilon^4] survives, the others carrying odd powers of ε\varepsilon; then E[ε4]=3\mathbb{E}[\varepsilon^4] = 3 and E[ε6]=15\mathbb{E}[\varepsilon^6] = 15.
  6. Var⁡(gS)=μ4σ2+18μ2+15σ2−4μ2=μ4σ2+14μ2+15σ2\operatorname{Var}(g_S) = \dfrac{\mu^4}{\sigma^2} + 18\mu^2 + 15\sigma^2 - 4\mu^2 = \dfrac{\mu^4}{\sigma^2} + 14\mu^2 + 15\sigma^2.Subtract the squared mean.
  7. Both estimators have mean 2μ2\mu; Var⁡(gR)=4σ2\operatorname{Var}(g_R) = 4\sigma^2 while Var⁡(gS)=μ4σ2+14μ2+15σ2\operatorname{Var}(g_S) = \dfrac{\mu^4}{\sigma^2} + 14\mu^2 + 15\sigma^2At μ=σ=1\mu = \sigma = 1 the variances are 44 and 3030; as σ→0\sigma \to 0 the score-function variance grows like μ4/σ2\mu^4/\sigma^2 while the gradient it estimates stays at 2μ2\mu. The reparameterised estimator uses f′f', so each sample reports which way to move; the score-function estimator uses only values of ff and must infer the direction from which samples scored higher. This gap is why VAEs reparameterise and why policy gradients, which cannot (the next page), need baselines.

Problem 7

A linear Gaussian model: p(z)=N(0,1)p(z) = \mathcal{N}(0, 1), p(x∣z)=N(wz+b,γ2)p(x \mid z) = \mathcal{N}(wz + b, \gamma^2) with w,b,γw, b, \gamma fixed, and q(z)=N(μ,σ2)q(z) = \mathcal{N}(\mu, \sigma^2). Show that

L(μ,σ2)=−12log⁡(2πγ2)−(x−wμ−b)2+w2σ22γ2−12(μ2+σ2−log⁡σ2−1).\mathcal{L}(\mu, \sigma^2) = -\tfrac12\log(2\pi\gamma^2) - \frac{(x - w\mu - b)^2 + w^2\sigma^2}{2\gamma^2} - \tfrac12\big(\mu^2 + \sigma^2 - \log\sigma^2 - 1\big).
  1. log⁡p(x∣z)=−12log⁡(2πγ2)−(x−wz−b)22γ2\log p(x \mid z) = -\tfrac12\log(2\pi\gamma^2) - \dfrac{(x - wz - b)^2}{2\gamma^2}.The Gaussian log-density, the multivariate-Gaussian page's Problem 1 with d=1d = 1.
  2. Eq[(x−wz−b)2]=Var⁡q(x−wz−b)+(Eq[x−wz−b])2=w2σ2+(x−wμ−b)2\mathbb{E}_q[(x - wz - b)^2] = \operatorname{Var}_q(x - wz - b) + \big(\mathbb{E}_q[x - wz - b]\big)^2 = w^2\sigma^2 + (x - w\mu - b)^2.The variance page's Problem 1: E[u2]=Var⁡(u)+(E[u])2\mathbb{E}[u^2] = \operatorname{Var}(u) + (\mathbb{E}[u])^2, with Var⁡(−wz)=w2σ2\operatorname{Var}(-wz) = w^2\sigma^2 and Eq[z]=μ\mathbb{E}_q[z] = \mu.
  3. Eq[log⁡p(x∣z)]=−12log⁡(2πγ2)−(x−wμ−b)2+w2σ22γ2\mathbb{E}_q[\log p(x \mid z)] = -\tfrac12\log(2\pi\gamma^2) - \dfrac{(x - w\mu - b)^2 + w^2\sigma^2}{2\gamma^2}.Linearity of Eq\mathbb{E}_q over step 1.
  4. KL⁡(q ∥ p(z))=12(μ2+σ2−log⁡σ2−1)\operatorname{KL}\big(q\,\|\,p(z)\big) = \tfrac12\big(\mu^2 + \sigma^2 - \log\sigma^2 - 1\big).Problem 4 with d=1d = 1, μ2=0\mu_2 = 0, σ2=1\sigma_2 = 1.
  5. L(μ,σ2)=−12log⁡(2πγ2)−(x−wμ−b)2+w2σ22γ2−12(μ2+σ2−log⁡σ2−1)\mathcal{L}(\mu, \sigma^2) = -\tfrac12\log(2\pi\gamma^2) - \dfrac{(x - w\mu - b)^2 + w^2\sigma^2}{2\gamma^2} - \tfrac12\big(\mu^2 + \sigma^2 - \log\sigma^2 - 1\big)Problem 3's first form, term by term. The reconstruction term penalises the mean error and also the spread, through w2σ2w^2\sigma^2: a wide qq decodes to a wide range of xx. The KL penalises a narrow qq through −log⁡σ2-\log\sigma^2. The two pull in opposite directions and Problem 8 finds the balance; because everything is closed-form here, the model is a test case for the sampled gradients of Problem 9.

Problem 8

Maximise Problem 7's L\mathcal{L} over μ\mu and σ2\sigma^2. Show that μ∗=w(x−b)γ2+w2\mu^* = \dfrac{w(x - b)}{\gamma^2 + w^2} and σ∗2=γ2γ2+w2\sigma^{*2} = \dfrac{\gamma^2}{\gamma^2 + w^2}, that q∗=N(μ∗,σ∗2)q^* = \mathcal{N}(\mu^*, \sigma^{*2}) is the exact posterior p(z∣x)p(z \mid x), and that the bound is tight: L∗=log⁡p(x)\mathcal{L}^* = \log p(x) with p(x)=N(x;b,w2+γ2)p(x) = \mathcal{N}(x; b, w^2 + \gamma^2).

  1. ∂L∂μ=w(x−wμ−b)γ2−μ=0\dfrac{\partial\mathcal{L}}{\partial\mu} = \dfrac{w(x - w\mu - b)}{\gamma^2} - \mu = 0 gives μ(1+w2γ2)=w(x−b)γ2\mu\Big(1 + \dfrac{w^2}{\gamma^2}\Big) = \dfrac{w(x - b)}{\gamma^2}, so μ∗=w(x−b)γ2+w2\mu^* = \dfrac{w(x - b)}{\gamma^2 + w^2}.Chain rule on the square, whose inner derivative in μ\mu is −w-w; the KL contributes −μ-\mu; multiply through by γ2\gamma^2.
  2. ∂L∂σ2=−w22γ2−12+12σ2=0\dfrac{\partial\mathcal{L}}{\partial\sigma^2} = -\dfrac{w^2}{2\gamma^2} - \dfrac12 + \dfrac{1}{2\sigma^2} = 0 gives 1σ2=1+w2γ2\dfrac{1}{\sigma^2} = 1 + \dfrac{w^2}{\gamma^2}, so σ∗2=γ2γ2+w2\sigma^{*2} = \dfrac{\gamma^2}{\gamma^2 + w^2}.Differentiate with σ2\sigma^2 as the variable: −log⁡σ2-\log\sigma^2 has derivative −1/σ2-1/\sigma^2.
  3. The stationary point is a maximum: ∂2L/∂μ2=−w2/γ2−1<0\partial^2\mathcal{L}/\partial\mu^2 = -w^2/\gamma^2 - 1 < 0, ∂2L/∂(σ2)2=−1/(2σ4)<0\partial^2\mathcal{L}/\partial(\sigma^2)^2 = -1/(2\sigma^4) < 0 and the mixed derivative is 00.A diagonal Hessian with negative entries is negative definite, the Hessians page's test.
  4. The exact posterior is p(z∣x)∝p(z) p(x∣z)∝exp⁡(−z22−(x−wz−b)22γ2)p(z \mid x) \propto p(z)\,p(x \mid z) \propto \exp\Big(-\dfrac{z^2}{2} - \dfrac{(x - wz - b)^2}{2\gamma^2}\Big), a Gaussian in zz with precision 1+w2/γ21 + w^2/\gamma^2 and mean w(x−b)/γ21+w2/γ2=w(x−b)γ2+w2\dfrac{w(x - b)/\gamma^2}{1 + w^2/\gamma^2} = \dfrac{w(x - b)}{\gamma^2 + w^2}: exactly (μ∗,σ∗2)(\mu^*, \sigma^{*2}).Collect the z2z^2 and zz terms in the exponent: the coefficient of −z2/2-z^2/2 is the precision and the coefficient of zz divided by the precision is the mean, as in the multivariate-Gaussian page's product of Gaussians.
  5. q∗=p(z∣x)q^* = p(z \mid x) makes the KL of Problem 2 zero, so L∗=log⁡p(x)\mathcal{L}^* = \log p(x); and x=wz+b+ηx = wz + b + \eta with z∼N(0,1)z \sim \mathcal{N}(0, 1) and η∼N(0,γ2)\eta \sim \mathcal{N}(0, \gamma^2) independent has mean bb and variance w2+γ2w^2 + \gamma^2.Problem 2; the multivariate-Gaussian page's affine image and the variance page's Problem 2 for the sum of independent terms.
  6. μ∗=w(x−b)γ2+w2\mu^* = \dfrac{w(x - b)}{\gamma^2 + w^2}, σ∗2=γ2γ2+w2\sigma^{*2} = \dfrac{\gamma^2}{\gamma^2 + w^2}; q∗q^* is the posterior p(z∣x)p(z \mid x), so the bound is tight: L∗=log⁡p(x)=log⁡N(x;b,w2+γ2)\mathcal{L}^* = \log p(x) = \log\mathcal{N}(x; b, w^2 + \gamma^2)The variational family contains the posterior, so the ELBO reaches the evidence; with a nonlinear decoder it does not, and the gap of Problem 2 remains. The posterior is narrower than the prior, σ∗2<1\sigma^{*2} < 1, by more when the signal-to-noise ratio w2/γ2w^2/\gamma^2 is large, and μ∗\mu^* shrinks the naive estimate (x−b)/w(x - b)/w towards 00 by the factor w2/(γ2+w2)w^2/(\gamma^2 + w^2). A VAE's encoder learns an amortised version of the map x↦(μ∗,σ∗)x \mapsto (\mu^*, \sigma^*).

Problem 9

Let the decoder give log⁡pθ(x∣z)=g(z)\log p_\theta(x \mid z) = g(z) for a differentiable gg, with q=N(μ,σ2)q = \mathcal{N}(\mu, \sigma^2) and prior N(0,1)\mathcal{N}(0, 1). The one-sample estimate of the ELBO is L^(μ,σ)=g(μ+σε)−12(μ2+σ2−log⁡σ2−1)\hat{\mathcal{L}}(\mu, \sigma) = g(\mu + \sigma\varepsilon) - \tfrac12\big(\mu^2 + \sigma^2 - \log\sigma^2 - 1\big) with ε\varepsilon drawn once. Compute ∂L^/∂μ\partial\hat{\mathcal{L}}/\partial\mu, ∂L^/∂σ\partial\hat{\mathcal{L}}/\partial\sigma and ∂L^/∂s\partial\hat{\mathcal{L}}/\partial s with s=log⁡σ2s = \log\sigma^2, and show that each has expectation equal to the corresponding gradient of L\mathcal{L}.

  1. z=μ+σεz = \mu + \sigma\varepsilon with ∂z/∂μ=1\partial z/\partial\mu = 1 and ∂z/∂σ=ε\partial z/\partial\sigma = \varepsilon.The reparameterisation; ε\varepsilon is a constant once drawn.
  2. ∂L^∂μ=g′(z)−μ\dfrac{\partial\hat{\mathcal{L}}}{\partial\mu} = g'(z) - \mu.Chain rule through zz; the KL term's derivative in μ\mu is μ\mu.
  3. ∂L^∂σ=g′(z) ε−σ+1σ\dfrac{\partial\hat{\mathcal{L}}}{\partial\sigma} = g'(z)\,\varepsilon - \sigma + \dfrac1\sigma.Chain rule through zz; 12(σ2−log⁡σ2)\tfrac12(\sigma^2 - \log\sigma^2) has derivative σ−1/σ\sigma - 1/\sigma.
  4. ∂L^∂s=σ2 g′(z) ε−12(σ2−1)\dfrac{\partial\hat{\mathcal{L}}}{\partial s} = \dfrac\sigma2\,g'(z)\,\varepsilon - \tfrac12(\sigma^2 - 1).σ=es/2\sigma = e^{s/2} gives ∂σ/∂s=σ/2\partial\sigma/\partial s = \sigma/2 for the first term; the KL in ss is 12(μ2+es−s−1)\tfrac12(\mu^2 + e^s - s - 1) with derivative 12(es−1)\tfrac12(e^s - 1), the entropy page's Problem 9.
  5. Eε[∂L^/∂μ]=E[g′(z)]−μ=∂μE[g(z)]−∂μKL⁡\mathbb{E}_\varepsilon\big[\partial\hat{\mathcal{L}}/\partial\mu\big] = \mathbb{E}[g'(z)] - \mu = \partial_\mu\mathbb{E}[g(z)] - \partial_\mu\operatorname{KL}, and likewise for σ\sigma and ss.Problem 5 with f=gf = g for the reconstruction term; the KL term contains no ε\varepsilon and is already exact.
  6. ∂L^∂μ=g′(z)−μ\dfrac{\partial\hat{\mathcal{L}}}{\partial\mu} = g'(z) - \mu; ∂L^∂σ=g′(z) ε−σ+1σ\dfrac{\partial\hat{\mathcal{L}}}{\partial\sigma} = g'(z)\,\varepsilon - \sigma + \dfrac1\sigma; ∂L^∂s=σ2 g′(z) ε−12(σ2−1)\dfrac{\partial\hat{\mathcal{L}}}{\partial s} = \dfrac\sigma2\,g'(z)\,\varepsilon - \tfrac12(\sigma^2 - 1), with z=μ+σεz = \mu + \sigma\varepsilon; each is an unbiased estimate of the gradient of L\mathcal{L}This is one VAE training step for one latent coordinate: the decoder's gradient g′(z)g'(z) at the sampled zz, passed back to μ\mu with weight 11 and to σ\sigma with weight ε\varepsilon, plus the closed-form KL gradient. Treating the sample as a constant (Mistake 2) deletes g′(z)εg'(z)\varepsilon and the encoder's variance stops learning from the data. With dd coordinates everything is per coordinate by Problem 4, and the check differentiates L^\hat{\mathcal{L}} numerically for a specific nonlinear gg with ε\varepsilon held fixed.

Problem 10

Instead of Problem 4's closed form, the KL can be estimated from one sample: K^=log⁡q(z)−log⁡p(z)\hat K = \log q(z) - \log p(z) with z=μ+σεz = \mu + \sigma\varepsilon, q=N(μ,σ2)q = \mathcal{N}(\mu, \sigma^2) and p=N(0,1)p = \mathcal{N}(0, 1). Show that K^=−log⁡σ+12μ2+μσε+12(σ2−1)ε2\hat K = -\log\sigma + \tfrac12\mu^2 + \mu\sigma\varepsilon + \tfrac12(\sigma^2 - 1)\varepsilon^2, that E[K^]=KL⁡(q ∥ p)\mathbb{E}[\hat K] = \operatorname{KL}(q\,\|\,p), and that Var⁡(K^)=μ2σ2+12(σ2−1)2\operatorname{Var}(\hat K) = \mu^2\sigma^2 + \tfrac12(\sigma^2 - 1)^2.

  1. log⁡q(z)=−12log⁡(2π)−log⁡σ−(z−μ)22σ2=−12log⁡(2π)−log⁡σ−12ε2\log q(z) = -\tfrac12\log(2\pi) - \log\sigma - \dfrac{(z - \mu)^2}{2\sigma^2} = -\tfrac12\log(2\pi) - \log\sigma - \tfrac12\varepsilon^2.(z−μ)/σ=ε(z - \mu)/\sigma = \varepsilon.
  2. log⁡p(z)=−12log⁡(2π)−12(μ+σε)2=−12log⁡(2π)−12μ2−μσε−12σ2ε2\log p(z) = -\tfrac12\log(2\pi) - \tfrac12(\mu + \sigma\varepsilon)^2 = -\tfrac12\log(2\pi) - \tfrac12\mu^2 - \mu\sigma\varepsilon - \tfrac12\sigma^2\varepsilon^2.Expand the square.
  3. K^=−log⁡σ−12ε2+12μ2+μσε+12σ2ε2=−log⁡σ+12μ2+μσε+12(σ2−1)ε2\hat K = -\log\sigma - \tfrac12\varepsilon^2 + \tfrac12\mu^2 + \mu\sigma\varepsilon + \tfrac12\sigma^2\varepsilon^2 = -\log\sigma + \tfrac12\mu^2 + \mu\sigma\varepsilon + \tfrac12(\sigma^2 - 1)\varepsilon^2.Subtract; the log⁡(2π)\log(2\pi) terms cancel and the ε2\varepsilon^2 terms combine.
  4. E[K^]=−log⁡σ+12μ2+0+12(σ2−1)=12(μ2+σ2−log⁡σ2−1)\mathbb{E}[\hat K] = -\log\sigma + \tfrac12\mu^2 + 0 + \tfrac12(\sigma^2 - 1) = \tfrac12\big(\mu^2 + \sigma^2 - \log\sigma^2 - 1\big).E[ε]=0\mathbb{E}[\varepsilon] = 0, E[ε2]=1\mathbb{E}[\varepsilon^2] = 1, and log⁡σ=12log⁡σ2\log\sigma = \tfrac12\log\sigma^2; this is Problem 4's closed form.
  5. Var⁡(K^)=μ2σ2Var⁡(ε)+14(σ2−1)2Var⁡(ε2)+2⋅μσ⋅12(σ2−1)Cov⁡(ε,ε2)=μ2σ2+12(σ2−1)2\operatorname{Var}(\hat K) = \mu^2\sigma^2\operatorname{Var}(\varepsilon) + \tfrac14(\sigma^2 - 1)^2\operatorname{Var}(\varepsilon^2) + 2\cdot\mu\sigma\cdot\tfrac12(\sigma^2 - 1)\operatorname{Cov}(\varepsilon, \varepsilon^2) = \mu^2\sigma^2 + \tfrac12(\sigma^2 - 1)^2.The variance page's Problem 2 on the two random terms of step 3; Var⁡(ε2)=E[ε4]−1=2\operatorname{Var}(\varepsilon^2) = \mathbb{E}[\varepsilon^4] - 1 = 2 and Cov⁡(ε,ε2)=E[ε3]−E[ε] E[ε2]=0\operatorname{Cov}(\varepsilon, \varepsilon^2) = \mathbb{E}[\varepsilon^3] - \mathbb{E}[\varepsilon]\,\mathbb{E}[\varepsilon^2] = 0.
  6. K^=−log⁡σ+12μ2+μσε+12(σ2−1)ε2\hat K = -\log\sigma + \tfrac12\mu^2 + \mu\sigma\varepsilon + \tfrac12(\sigma^2 - 1)\varepsilon^2; E[K^]=KL⁡(q ∥ p)\mathbb{E}[\hat K] = \operatorname{KL}(q\,\|\,p); Var⁡(K^)=μ2σ2+12(σ2−1)2\operatorname{Var}(\hat K) = \mu^2\sigma^2 + \tfrac12(\sigma^2 - 1)^2The estimator is unbiased, so a VAE trained with it is right on average, but its noise is pure cost: at μ=2\mu = 2, σ=1\sigma = 1 the variance is 44 for a KL of 22. The closed form removes it entirely, which is why implementations compute the KL analytically whenever qq and the prior are Gaussian and sample only the reconstruction term. At q=pq = p the estimator is identically 00, so near the prior it is quiet; sampling from the prior instead of qq (Mistake 3) breaks even the unbiasedness.

Where this goes wrong

1. The KL term written the other way round

KL is not symmetric, and the regulariser is often described as the distance between qq and the prior.

  1. L=Eq[log⁡p(x∣z)]−KL⁡(q(z) ∥ p(z))\mathcal{L} = \mathbb{E}_q[\log p(x \mid z)] - \operatorname{KL}\big(q(z)\,\|\,p(z)\big)Right so far: Problem 3.
  2. “KL is the distance from qq to the prior, so the order does not matter.”The habit that causes the mistake: treating KL as a symmetric distance, which it is not (the entropy page's Problem 6).
  3. L=Eq[log⁡p(x∣z)]−KL⁡(p(z) ∥ q(z))\mathcal{L} = \mathbb{E}_q[\log p(x \mid z)] - \operatorname{KL}\big(p(z)\,\|\,q(z)\big)For q=N(μ,σ2)q = \mathcal{N}(\mu, \sigma^2) and p=N(0,1)p = \mathcal{N}(0, 1) the reversed term is log⁡σ+1+μ22σ2−12\log\sigma + \dfrac{1 + \mu^2}{2\sigma^2} - \dfrac12, which at μ=0\mu = 0, σ=2\sigma = 2 is 0.3180.318 where the correct 12(σ2−log⁡σ2−1)\tfrac12(\sigma^2 - \log\sigma^2 - 1) is 0.8070.807. The result is no longer a lower bound on log⁡p(x)\log p(x) and Problem 2's identity fails. Its gradient in σ\sigma is different in kind: the reversed KL punishes a small σ\sigma through 1/σ21/\sigma^2 and a large one only logarithmically, so it pushes qq to be wide, the mode-covering behaviour of the entropy page's Problem 7 rather than the mode-seeking one the ELBO has.

2. Treating the sample as a constant

After sampling, zz is a tensor of numbers, and dist.sample() returns exactly that, with no gradient path back to μ\mu or σ\sigma.

  1. L^=g(z)−12(μ2+σ2−log⁡σ2−1)\hat{\mathcal{L}} = g(z) - \tfrac12(\mu^2 + \sigma^2 - \log\sigma^2 - 1) with z=μ+σεz = \mu + \sigma\varepsilonRight so far: Problem 9.
  2. “zz has been sampled, so it is a number now; only the KL term depends on σ\sigma.”The shortcut that causes the mistake: sampling with sample() instead of rsample(), which detaches zz from its parameters.
  3. ∂L^∂σ=−σ+1σ\dfrac{\partial\hat{\mathcal{L}}}{\partial\sigma} = -\sigma + \dfrac1\sigmaThe reconstruction term's dependence on σ\sigma through z=μ+σεz = \mu + \sigma\varepsilon is gone, and with it g′(z)εg'(z)\varepsilon (Problem 9); likewise ∂L^/∂μ\partial\hat{\mathcal{L}}/\partial\mu loses g′(z)g'(z). What remains is the KL's gradient alone, which is zero at σ=1\sigma = 1, μ=0\mu = 0 and pushes the encoder there whatever the data say: qq collapses onto the prior and the decoder is trained on noise. The reparameterisation of Problem 5 exists to keep that path; zz must be written as a function of (μ,σ)(\mu, \sigma) and ε\varepsilon before anything is differentiated.

3. Monte Carlo KL averaged over samples from the prior

The KL is an integral of log⁡q−log⁡p\log q - \log p, and the prior N(0,I)\mathcal{N}(0, I) is the easiest distribution in sight to sample.

  1. KL⁡(q ∥ p)=Eq[log⁡q(z)−log⁡p(z)]\operatorname{KL}(q\,\|\,p) = \mathbb{E}_q[\log q(z) - \log p(z)]Right so far: the definition.
  2. “Draw zz from the prior, which needs no parameters, and average log⁡q−log⁡p\log q - \log p.”The shortcut that causes the mistake: an expectation under qq estimated with samples from pp, which is a different integral.
  3. KL⁡(q ∥ p)≈1K∑k=1K(log⁡q(zk)−log⁡p(zk))\operatorname{KL}(q\,\|\,p) \approx \dfrac1K\sum_{k=1}^{K}\big(\log q(z_k) - \log p(z_k)\big) with zk∼pz_k \sim pIts expectation is Ep[log⁡q−log⁡p]=−KL⁡(p ∥ q)\mathbb{E}_p[\log q - \log p] = -\operatorname{KL}(p\,\|\,q), which is never positive, so the estimate is negative on average although a KL is never negative: at μ=0\mu = 0, σ=2\sigma = 2 it averages −0.318-0.318 where the KL is 0.8070.807. The estimator of Problem 10 draws z=μ+σεz = \mu + \sigma\varepsilon from qq and is unbiased; the samples must come from the distribution the expectation is under.

4. Jensen's inequality applied in the wrong direction

The inequality says that the log of an average and the average of the logs differ, and which is larger is easy to misremember.

  1. log⁡p(x)=log⁡Eq[p(x,z)/q(z)]\log p(x) = \log\mathbb{E}_q\big[p(x, z)/q(z)\big]Right so far: Problem 1, step 1.
  2. “Move the log inside: log⁡E[w]≤E[log⁡w]\log\mathbb{E}[w] \le \mathbb{E}[\log w].”The habit that causes the mistake: the direction for a convex function such as the square, E[w2]≥(E[w])2\mathbb{E}[w^2] \ge (\mathbb{E}[w])^2, applied to the concave logarithm.
  3. log⁡p(x)≤Eq[log⁡p(x,z)−log⁡q(z)]\log p(x) \le \mathbb{E}_q\big[\log p(x, z) - \log q(z)\big]For the concave logarithm, log⁡E[w]≥E[log⁡w]\log\mathbb{E}[w] \ge \mathbb{E}[\log w], so the ELBO is a lower bound. The check is Problem 2: log⁡p(x)−L(q)=KL⁡(q ∥ p(z∣x))≥0\log p(x) - \mathcal{L}(q) = \operatorname{KL}\big(q\,\|\,p(z \mid x)\big) \ge 0, and on the discrete model in the check L\mathcal{L} is strictly below log⁡p(x)\log p(x) for every qq that is not the posterior. Were the line true, maximising L\mathcal{L} would raise an upper bound, which says nothing about log⁡p(x)\log p(x).

5. Reparameterising with the variance in place of the standard deviation

The encoder outputs s=log⁡σ2s = \log\sigma^2, and the sampling line needs a scale.

  1. z=μ+σ⊙εz = \mu + \sigma\odot\varepsilon with σj=esj/2\sigma_j = e^{s_j/2}Right so far: Problem 5's reparameterisation.
  2. “The encoder's scale is ese^{s}, so multiply ε\varepsilon by it.”The shortcut that causes the mistake: exponentiating the log-variance gives the variance, and the sample needs the standard deviation.
  3. z=μ+es⊙εz = \mu + e^{s}\odot\varepsilonThen z∼N(μ,diag⁡(e2s))=N(μ,diag⁡(σ4))z \sim \mathcal{N}\big(\mu, \operatorname{diag}(e^{2s})\big) = \mathcal{N}\big(\mu, \operatorname{diag}(\sigma^4)\big): the sample's standard deviation is σ2\sigma^2, too wide when σ>1\sigma > 1 and too narrow when σ<1\sigma < 1. The KL term, computed from ss for N(μ,σ2)\mathcal{N}(\mu, \sigma^2) (Problem 4), then regularises a different distribution from the one that was sampled, and the two halves of the ELBO describe two different qq's. The scale is es/2e^{s/2}, exp(0.5 * logvar) in code; the multivariate-Gaussian page's x=μ+Σεx = \mu + \Sigma\varepsilon is the full-covariance form of the same error.

Print this set: elbo-and-the-reparameterisation-trick.pdf (problems, answers, and worked solutions on separate pages).