import nbformat as nbf nb = nbf.v4.new_notebook(); C=[] md = lambda s: C.append(nbf.v4.new_markdown_cell(s)) code = lambda s: C.append(nbf.v4.new_code_cell(s)) md("""# VAE lecture demo (Week 10) — CS6140 Companion to the lecture note **VAE_claude_0929.pdf**. Each cell checks one board computation numerically, and the last one trains a tiny autoencoder and a tiny VAE side by side. Data here: sklearn's **8×8 digits** with a **Gaussian** decoder (squared error); HW5 Problem 3 uses 28×28 MNIST with a Bernoulli decoder (binary cross-entropy) — nothing here is its solution. | Cell | Board segment | What it checks | |---|---|---| | 1 | setup | imports, device (tiny models: CPU) | | 2 | KL from scratch | discrete KL, asymmetry, cross-entropy = entropy + KL | | 3 | KL of Gaussians | closed form vs Monte-Carlo; what the KL penalty "costs" | | 4 | ELBO on a 2-state latent | ELBO + KL(q‖posterior) = log p(x) for any q; q = posterior is EM's E-step | | 5 | linear decoder (probabilistic PCA) | the exact posterior is Gaussian with a *linear* mean: a one-layer encoder | | 6 | reparameterization | `sample()` has no gradient, `rsample()` = μ + σε does; gradient of E[z²] | | 7 | AE vs VAE | same network ± KL term: latent layout and decoding a grid of z | Total run time: well under a minute on a laptop CPU.""") code("""import math, numpy as np, torch, torch.nn as nn, matplotlib.pyplot as plt, os from sklearn.datasets import load_digits torch.set_num_threads(4) def get_device(): if torch.cuda.is_available(): return torch.device('cuda') if torch.backends.mps.is_available(): return torch.device('mps') return torch.device('cpu') print('best available device:', get_device()) DEVICE = torch.device('cpu') # tiny models: the CPU is fastest FIG = 'figures_0929'; os.makedirs(FIG, exist_ok=True)""") md("""## 2. KL divergence from scratch (discrete) $\\mathrm{KL}(p\\|q)=\\sum_x p(x)\\log\\frac{p(x)}{q(x)}$ — the extra "surprise" you pay for believing $q$ when the truth is $p$.""") code("""p = np.array([.5, .25, .25]); q = np.array([1/3, 1/3, 1/3]) KL = lambda a, b: np.sum(a*np.log(a/b)) H = -np.sum(p*np.log(p)); CE = -np.sum(p*np.log(q)) print(f'KL(p||q) = {KL(p,q):.4f} KL(q||p) = {KL(q,p):.4f} (not symmetric; both >= 0)') print(f'entropy H(p) = {H:.4f}, cross-entropy H(p,q) = {CE:.4f} = H(p) + KL(p||q) = {H+KL(p,q):.4f}') print('=> minimizing cross-entropy (= the softmax/logistic loss, = maximizing likelihood) over q <=> minimizing KL(p_data || q)')""") md("""## 3. KL between Gaussians: closed form vs. Monte-Carlo $\\mathrm{KL}(\\mathcal N(\\mu,\\sigma^2)\\,\\|\\,\\mathcal N(0,1))=\\tfrac12(\\mu^2+\\sigma^2-1-\\log\\sigma^2)$.""") code("""kl_gauss = lambda mu, s: 0.5*(mu**2 + s**2 - 1 - np.log(s**2)) rng = np.random.default_rng(0) print(' mu sigma closed form Monte-Carlo (10^6 samples of log q(z) - log p(z))') for mu, s in [(0, 1), (1, 1), (0, 2), (0, .5), (1, .5), (0, .1)]: z = mu + s*rng.standard_normal(10**6) mc = np.mean((-0.5*np.log(2*np.pi*s*s) - (z-mu)**2/(2*s*s)) - (-0.5*np.log(2*np.pi) - z**2/2)) print(f'{mu:5.1f} {s:6.2f} {kl_gauss(mu, s):10.4f} {mc:10.4f}') fig, ax = plt.subplots(1, 2, figsize=(9, 2.8)) m = np.linspace(-3, 3, 200); ax[0].plot(m, kl_gauss(m, 1.0)); ax[0].set_xlabel('mu (sigma = 1)'); ax[0].set_title('moving the mean: KL = mu^2 / 2') s = np.linspace(0.05, 3, 300); ax[1].plot(s, kl_gauss(0, s)); ax[1].set_xlabel('sigma (mu = 0)'); ax[1].set_title('shrinking sigma -> 0 costs without bound') for a in ax: a.set_ylabel('KL to N(0,1)'); a.grid(alpha=.3) plt.tight_layout(); plt.savefig(f'{FIG}/fig_kl_gauss.png', dpi=130); plt.show()""") md("""## 4. The ELBO on a two-state latent variable (HW4's mixture, with K = 2) Prior $p(z)=(0.5,0.5)$; for the observed point, $p(x\\mid z=1)=0.2$, $p(x\\mid z=2)=0.6$. So $p(x)=0.4$ and the posterior is $(0.25, 0.75)$.""") code("""pz = np.array([.5, .5]); px_z = np.array([.2, .6]) px = pz @ px_z; post = pz*px_z/px print(f'log p(x) = {np.log(px):.4f}, posterior = {post}') print(' q(z) recon E_q[log p(x|z)] KL(q||prior) ELBO gap KL(q||posterior)') for qz in [np.array([.5, .5]), np.array([.1, .9]), post]: recon = qz @ np.log(px_z); klp = KL(qz, pz); elbo = recon - klp print(f'{str(qz.round(2)):14s} {recon:12.4f} {klp:10.4f} {elbo:8.4f} {np.log(px)-elbo:6.4f} {KL(qz, post):8.4f}') print('q = posterior closes the gap: that is exactly what the EM E-step (responsibilities) computes.')""") md("""## 5. A linear decoder: everything is computable (probabilistic PCA) $z\\sim\\mathcal N(0,1)$, $x = wz+\\varepsilon$, $\\varepsilon\\sim\\mathcal N(0,s^2I)$ with $w=(2,1)$, $s^2=1$. Bayes' rule gives $p(z\\mid x)=\\mathcal N\\big(w^\\top x/M,\\ s^2/M\\big)$, $M=w^\\top w+s^2=6$. Check it by brute force on a grid for $x=(3,1)$.""") code("""w = np.array([2., 1.]); s2 = 1.; x = np.array([3., 1.]); M = w@w + s2 zz = np.linspace(-6, 6, 200001) logjoint = -zz**2/2 - ((x[None, :] - np.outer(zz, w))**2).sum(1)/(2*s2) # log p(z) + log p(x|z), up to constants pj = np.exp(logjoint - logjoint.max()); pj /= pj.sum() mean = (zz*pj).sum(); var = ((zz-mean)**2*pj).sum() print(f'formula: mean = w.x/M = {w@x/M:.4f}, var = s^2/M = {s2/M:.4f} grid: mean = {mean:.4f}, var = {var:.4f}') print(f'the exact "encoder" is linear: mu(x) = [{w[0]/M:.3f}, {w[1]/M:.3f}] . x — one nn.Linear(2, 1) layer.') print(f'as s^2 -> 0 the mean becomes w.x/|w|^2 = {w@x/(w@w):.2f}: the projection onto w, i.e. PCA (HW3)') print('marginal: p(x) = N(0, w w^T + s^2 I) =\\n', np.outer(w, w) + s2*np.eye(2))""") md("""## 6. The reparameterization trick We need $\\frac{\\partial}{\\partial\\mu}$ and $\\frac{\\partial}{\\partial\\sigma}$ of $\\mathbb E_{z\\sim\\mathcal N(\\mu,\\sigma^2)}[f(z)]$. For $f(z)=z^2$ the answer is known: $\\mathbb E[z^2]=\\mu^2+\\sigma^2$, so the gradients are $2\\mu$ and $2\\sigma$.""") code("""mu = torch.tensor(1.0, requires_grad=True); sigma = torch.tensor(0.5, requires_grad=True) dist = torch.distributions.Normal(mu, sigma) z_plain = dist.sample((4,)); print('sample(): requires_grad =', z_plain.requires_grad, ' -> backprop stops here') z_rep = dist.rsample((4,)); print('rsample(): requires_grad =', z_rep.requires_grad, ' (rsample computes mu + sigma*eps)') eps = torch.tensor([-1., 1.]) # the two board samples z = mu + sigma*eps; (z**2).mean().backward() print(f'2 samples eps=(-1,+1): d/dmu = {mu.grad.item():.3f} (exact 2mu = 2), d/dsigma = {sigma.grad.item():.3f} (exact 2sigma = 1)') mu.grad = None; sigma.grad = None torch.manual_seed(0); z = mu + sigma*torch.randn(100000); (z**2).mean().backward() print(f'100k samples: d/dmu = {mu.grad.item():.3f}, d/dsigma = {sigma.grad.item():.3f}')""") md("""## 7. Autoencoder vs. VAE: same network, one extra term Encoder 64 → 128 → (2 means, 2 log-variances); decoder 2 → 128 → 64 with sigmoid outputs. The **AE** uses the mean only and minimizes reconstruction error. The **VAE** samples $z=\\mu+\\sigma\\epsilon$ and adds the KL term. Reconstruction term: Gaussian likelihood with $\\sigma_x=0.25$, i.e. $\\|x-\\hat x\\|^2/(2\\cdot0.25^2)$ — Week 1's "squared error = Gaussian log-likelihood". What to look for: with 2 latent dimensions and small 8×8 images, *both* decoders produce digit-like pictures somewhere. The difference is **where the codes are**: the AE's code scale and location are arbitrary (whatever reconstruction happened to like), so "pick a random z" has no answer; the VAE's KL term pins the codes to $\\mathcal N(0,I)$, so sampling the prior *is* the answer.""") code("""X, y = load_digits(return_X_y=True); X = torch.tensor(X/16., dtype=torch.float32) class AEorVAE(nn.Module): def __init__(self, variational): super().__init__(); self.variational = variational self.trunk = nn.Sequential(nn.Linear(64, 128), nn.ReLU()) self.head_mu, self.head_logvar = nn.Linear(128, 2), nn.Linear(128, 2) # two heads, one trunk self.decoder = nn.Sequential(nn.Linear(2, 128), nn.ReLU(), nn.Linear(128, 64), nn.Sigmoid()) def forward(self, x): h = self.trunk(x); mu = self.head_mu(h) if not self.variational: return self.decoder(mu), mu, None logvar = self.head_logvar(h) z = mu + torch.exp(0.5*logvar)*torch.randn_like(mu) # reparameterization return self.decoder(z), mu, logvar def fit(variational, epochs=1500, sx=0.25): torch.manual_seed(0); net = AEorVAE(variational); opt = torch.optim.Adam(net.parameters(), lr=3e-3) for ep in range(epochs): # full batch: only 1797 images xhat, mu, logvar = net(X) recon = ((X - xhat)**2).sum(1).mean()/(2*sx**2) kl = 0.5*(mu**2 + logvar.exp() - 1 - logvar).sum(1).mean() if variational else torch.tensor(0.) loss = recon + kl; opt.zero_grad(); loss.backward(); opt.step() print(f"{'VAE' if variational else 'AE ':3s}: reconstruction term {recon.item():6.1f}, KL term {kl.item():5.2f}") return net ae, vae = fit(False), fit(True)""") code("""with torch.no_grad(): codes = {'autoencoder': ae(X)[1].numpy(), 'VAE (means)': vae(X)[1].numpy()} fig, ax = plt.subplots(1, 2, figsize=(10, 4.2)) for a, (name, c) in zip(ax, codes.items()): sc = a.scatter(c[:, 0], c[:, 1], c=y, cmap='tab10', s=6); a.set_title(f'{name}: code std = {c.std(0).round(1)}') a.set_aspect('equal', 'datalim') fig.colorbar(sc, ax=ax, ticks=range(10), label='digit'); plt.savefig(f'{FIG}/fig_ae_vs_vae_latent.png', dpi=130, bbox_inches='tight'); plt.show()""") code("""def tile(imgs, n_cols, pad=1): # lay out 8x8 images on a white canvas with a 1-pixel gap between tiles n_rows = int(np.ceil(len(imgs)/n_cols)); canvas = np.zeros((n_rows*(8+pad)+pad, n_cols*(8+pad)+pad)) for k, im in enumerate(imgs): r, c = divmod(k, n_cols); canvas[pad+r*(8+pad):pad+r*(8+pad)+8, pad+c*(8+pad):pad+c*(8+pad)+8] = im.reshape(8, 8) return canvas def decode_grid(net, lim_x, lim_y, n=10): gx, gy = np.linspace(*lim_x, n), np.linspace(*lim_y, n) Z = torch.tensor([[a, b] for b in gy[::-1] for a in gx], dtype=torch.float32) with torch.no_grad(): return tile(net.decoder(Z).numpy(), n) c = codes['autoencoder']; lo, hi = np.percentile(c, 1, axis=0), np.percentile(c, 99, axis=0) fig, ax = plt.subplots(1, 2, figsize=(10, 5.2)) ax[0].imshow(decode_grid(ae, (lo[0], hi[0]), (lo[1], hi[1])), cmap='gray_r') ax[0].set_title('AE: even grid over the box holding 98% of its codes (about -30..30)\\n(the decoder works, but you had to look up where the codes are)', fontsize=9) ax[1].imshow(decode_grid(vae, (-2.5, 2.5), (-2.5, 2.5)), cmap='gray_r') ax[1].set_title('VAE: even grid over [-2.5, 2.5]^2, i.e. where N(0, I) puts its mass\\n(all ten digits; neighbouring cells change gradually)', fontsize=9) for a in ax: a.axis('off') plt.tight_layout(); plt.savefig(f'{FIG}/fig_decode_grids.png', dpi=130); plt.show()""") code("""torch.manual_seed(1); zs = torch.randn(12, 2) with torch.no_grad(): rows = [ae.decoder(zs).numpy(), vae.decoder(zs).numpy()] fig, ax = plt.subplots(2, 1, figsize=(9, 2.4)) for a, r, t in zip(ax, rows, ['AE decoder on z ~ N(0, I): these z cover only the centre of its code space (code std about 8), so only the few classes clustered near 0 appear', 'VAE decoder on the same z: N(0, I) is where the KL term put the codes, so the samples spread over many classes (this is "generation")']): a.imshow(tile(r, 12), cmap='gray_r'); a.set_title(t, fontsize=8.5, loc='left'); a.axis('off') plt.tight_layout(); plt.savefig(f'{FIG}/fig_prior_samples.png', dpi=130); plt.show()""") nb['cells'] = C nb.metadata['kernelspec'] = {"name": "python3", "display_name": "Python 3", "language": "python"} nbf.write(nb, '/Users/vip/Dropbox/CS6140/3_generative_models/lecture_notes/VAE/vae_lecture_demo.ipynb')