|
119 | 119 | }, |
120 | 120 | { |
121 | 121 | "cell_type": "code", |
122 | | - "execution_count": 3, |
| 122 | + "execution_count": null, |
123 | 123 | "id": "9638499c", |
124 | 124 | "metadata": {}, |
125 | 125 | "outputs": [ |
|
151 | 151 | " \"\"\"Convolutional Denoising Autoencoder\"\"\"\n", |
152 | 152 | " def __init__(self):\n", |
153 | 153 | " super().__init__()\n", |
154 | | - " # Encoder: 28×28 → 14×14 → 7×7\n", |
| 154 | + " # Encoder: 28x28 → 14x14 → 7x7\n", |
155 | 155 | " self.encoder = nn.Sequential(\n", |
156 | | - " nn.Conv2d(1, 16, 3, stride=2, padding=1), # → 14×14×16\n", |
| 156 | + " nn.Conv2d(1, 16, 3, stride=2, padding=1), # → 14x14x16\n", |
157 | 157 | " nn.ReLU(),\n", |
158 | | - " nn.Conv2d(16, 32, 3, stride=2, padding=1), # → 7×7×32\n", |
| 158 | + " nn.Conv2d(16, 32, 3, stride=2, padding=1), # → 7x7x32\n", |
159 | 159 | " nn.ReLU(),\n", |
160 | 160 | " )\n", |
161 | | - " # Decoder: 7×7 → 14×14 → 28×28\n", |
| 161 | + " # Decoder: 7x7 → 14x14 → 28x28\n", |
162 | 162 | " self.decoder = nn.Sequential(\n", |
163 | | - " nn.ConvTranspose2d(32, 16, 4, stride=2, padding=1), # → 14×14×16\n", |
| 163 | + " nn.ConvTranspose2d(32, 16, 4, stride=2, padding=1), # → 14x14x16\n", |
164 | 164 | " nn.ReLU(),\n", |
165 | | - " nn.ConvTranspose2d(16, 1, 4, stride=2, padding=1), # → 28×28×1\n", |
| 165 | + " nn.ConvTranspose2d(16, 1, 4, stride=2, padding=1), # → 28x28x1\n", |
166 | 166 | " nn.Sigmoid() # Output in [0, 1]\n", |
167 | 167 | " )\n", |
168 | | - " \n", |
| 168 | + "\n", |
169 | 169 | " def forward(self, x):\n", |
170 | 170 | " z = self.encoder(x)\n", |
171 | 171 | " x_recon = self.decoder(z)\n", |
|
250 | 250 | " x_clean = x_clean.to(device)\n", |
251 | 251 | " x_noisy = add_noise(x_clean, sigma)\n", |
252 | 252 | " x_recon = model(x_noisy)\n", |
253 | | - " \n", |
| 253 | + "\n", |
254 | 254 | " n = 8\n", |
255 | 255 | " fig, axes = plt.subplots(3, n, figsize=(1.5*n, 4.5))\n", |
256 | 256 | " for i in range(n):\n", |
|
260 | 260 | " axes[1, i].axis('off')\n", |
261 | 261 | " axes[2, i].imshow(x_recon[i, 0].cpu(), cmap='gray')\n", |
262 | 262 | " axes[2, i].axis('off')\n", |
263 | | - " \n", |
| 263 | + "\n", |
264 | 264 | " axes[0, 0].set_ylabel('Clean', fontsize=12)\n", |
265 | 265 | " axes[1, 0].set_ylabel('Noisy', fontsize=12)\n", |
266 | 266 | " axes[2, 0].set_ylabel('Reconstructed', fontsize=12)\n", |
|
353 | 353 | " except StopIteration:\n", |
354 | 354 | " train_iter = iter(train_loader)\n", |
355 | 355 | " x_clean, _ = next(train_iter)\n", |
356 | | - " \n", |
| 356 | + "\n", |
357 | 357 | " x_clean = x_clean.to(device)\n", |
358 | 358 | " x_noisy = add_noise(x_clean, sigma=noise_sigma)\n", |
359 | | - " \n", |
| 359 | + "\n", |
360 | 360 | " # Forward pass\n", |
361 | 361 | " x_recon = dae(x_noisy)\n", |
362 | 362 | " loss = F.mse_loss(x_recon, x_clean)\n", |
363 | | - " \n", |
| 363 | + "\n", |
364 | 364 | " # Backward pass\n", |
365 | 365 | " optimizer_dae.zero_grad()\n", |
366 | 366 | " loss.backward()\n", |
367 | 367 | " optimizer_dae.step()\n", |
368 | | - " \n", |
| 368 | + "\n", |
369 | 369 | " losses.append(loss.item())\n", |
370 | | - " \n", |
| 370 | + "\n", |
371 | 371 | " if step % 200 == 0:\n", |
372 | 372 | " print(f\"Step {step}/{num_steps} | Loss: {loss.item():.4f}\")\n", |
373 | 373 | "\n", |
|
554 | 554 | " def __init__(self, latent_dim=8):\n", |
555 | 555 | " super().__init__()\n", |
556 | 556 | " self.latent_dim = latent_dim\n", |
557 | | - " \n", |
| 557 | + "\n", |
558 | 558 | " # Encoder\n", |
559 | 559 | " self.fc1 = nn.Linear(28*28, 256)\n", |
560 | 560 | " self.fc_mu = nn.Linear(256, latent_dim)\n", |
561 | 561 | " self.fc_logvar = nn.Linear(256, latent_dim)\n", |
562 | | - " \n", |
| 562 | + "\n", |
563 | 563 | " # Decoder\n", |
564 | 564 | " self.fc2 = nn.Linear(latent_dim, 256)\n", |
565 | 565 | " self.fc3 = nn.Linear(256, 28*28)\n", |
566 | | - " \n", |
| 566 | + "\n", |
567 | 567 | " def encode(self, x):\n", |
568 | 568 | " \"\"\"Encode input to latent distribution parameters.\"\"\"\n", |
569 | 569 | " h = F.relu(self.fc1(x))\n", |
570 | 570 | " mu = self.fc_mu(h)\n", |
571 | 571 | " logvar = self.fc_logvar(h)\n", |
572 | 572 | " return mu, logvar\n", |
573 | | - " \n", |
| 573 | + "\n", |
574 | 574 | " def reparameterize(self, mu, logvar):\n", |
575 | 575 | " \"\"\"Sample z from N(mu, sigma^2) using reparameterization trick.\"\"\"\n", |
576 | 576 | " std = torch.exp(0.5 * logvar)\n", |
577 | 577 | " eps = torch.randn_like(std)\n", |
578 | 578 | " z = mu + eps * std\n", |
579 | 579 | " return z\n", |
580 | | - " \n", |
| 580 | + "\n", |
581 | 581 | " def decode(self, z):\n", |
582 | 582 | " \"\"\"Decode latent code to reconstruction.\"\"\"\n", |
583 | 583 | " h = F.relu(self.fc2(z))\n", |
584 | 584 | " x_recon = torch.sigmoid(self.fc3(h))\n", |
585 | 585 | " return x_recon\n", |
586 | | - " \n", |
| 586 | + "\n", |
587 | 587 | " def forward(self, x):\n", |
588 | 588 | " mu, logvar = self.encode(x)\n", |
589 | 589 | " z = self.reparameterize(mu, logvar)\n", |
|
634 | 634 | "def vae_loss(x_recon, x, mu, logvar, beta=1.0):\n", |
635 | 635 | " \"\"\"Compute VAE loss = reconstruction + beta * KL divergence.\"\"\"\n", |
636 | 636 | " batch_size = x.size(0)\n", |
637 | | - " \n", |
| 637 | + "\n", |
638 | 638 | " # Reconstruction loss (binary cross-entropy)\n", |
639 | 639 | " recon_loss = F.binary_cross_entropy(x_recon, x, reduction='sum') / batch_size\n", |
640 | | - " \n", |
| 640 | + "\n", |
641 | 641 | " # KL divergence: KL(N(mu, sigma^2) || N(0, 1))\n", |
642 | 642 | " # = 0.5 * sum(exp(logvar) + mu^2 - 1 - logvar)\n", |
643 | 643 | " kl_loss = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp()) / batch_size\n", |
644 | | - " \n", |
| 644 | + "\n", |
645 | 645 | " total_loss = recon_loss + beta * kl_loss\n", |
646 | | - " \n", |
| 646 | + "\n", |
647 | 647 | " return total_loss, recon_loss, kl_loss" |
648 | 648 | ] |
649 | 649 | }, |
|
668 | 668 | " model.eval()\n", |
669 | 669 | " z = torch.randn(n, model.latent_dim).to(device)\n", |
670 | 670 | " x_samples = model.decode(z).view(n, 1, 28, 28).cpu()\n", |
671 | | - " \n", |
| 671 | + "\n", |
672 | 672 | " fig, axes = plt.subplots(4, 4, figsize=(6, 6))\n", |
673 | 673 | " for i, ax in enumerate(axes.flat):\n", |
674 | 674 | " ax.imshow(x_samples[i, 0], cmap='gray')\n", |
|
683 | 683 | " model.eval()\n", |
684 | 684 | " x1 = x1.view(1, -1).to(device)\n", |
685 | 685 | " x2 = x2.view(1, -1).to(device)\n", |
686 | | - " \n", |
| 686 | + "\n", |
687 | 687 | " # Encode to latent space (use mean for cleaner interpolation)\n", |
688 | 688 | " mu1, _ = model.encode(x1)\n", |
689 | 689 | " mu2, _ = model.encode(x2)\n", |
690 | | - " \n", |
| 690 | + "\n", |
691 | 691 | " # Linear interpolation in latent space\n", |
692 | 692 | " alphas = torch.linspace(0, 1, steps)\n", |
693 | 693 | " latents = [(1 - alpha) * mu1 + alpha * mu2 for alpha in alphas]\n", |
694 | | - " \n", |
| 694 | + "\n", |
695 | 695 | " # Decode interpolated latents\n", |
696 | 696 | " images = [model.decode(z).view(28, 28).cpu() for z in latents]\n", |
697 | | - " \n", |
| 697 | + "\n", |
698 | 698 | " fig, axes = plt.subplots(1, steps, figsize=(1.5*steps, 2))\n", |
699 | 699 | " for i, ax in enumerate(axes):\n", |
700 | 700 | " ax.imshow(images[i], cmap='gray')\n", |
|
712 | 712 | " x_flat = x.view(x.size(0), -1)\n", |
713 | 713 | " x_recon, _, _ = model(x_flat)\n", |
714 | 714 | " x_recon = x_recon.view(-1, 1, 28, 28)\n", |
715 | | - " \n", |
| 715 | + "\n", |
716 | 716 | " fig, axes = plt.subplots(2, n, figsize=(1.5*n, 3))\n", |
717 | 717 | " for i in range(n):\n", |
718 | 718 | " axes[0, i].imshow(x[i, 0].cpu(), cmap='gray')\n", |
|
776 | 776 | " except StopIteration:\n", |
777 | 777 | " train_iter = iter(train_loader)\n", |
778 | 778 | " x, _ = next(train_iter)\n", |
779 | | - " \n", |
| 779 | + "\n", |
780 | 780 | " x = x.view(x.size(0), -1).to(device)\n", |
781 | | - " \n", |
| 781 | + "\n", |
782 | 782 | " # Forward pass\n", |
783 | 783 | " x_recon, mu, logvar = vae(x)\n", |
784 | 784 | " loss, recon, kl = vae_loss(x_recon, x, mu, logvar, beta=beta)\n", |
785 | | - " \n", |
| 785 | + "\n", |
786 | 786 | " # Backward pass\n", |
787 | 787 | " optimizer_vae.zero_grad()\n", |
788 | 788 | " loss.backward()\n", |
789 | 789 | " optimizer_vae.step()\n", |
790 | | - " \n", |
| 790 | + "\n", |
791 | 791 | " total_losses.append(loss.item())\n", |
792 | 792 | " recon_losses.append(recon.item())\n", |
793 | 793 | " kl_losses.append(kl.item())\n", |
794 | | - " \n", |
| 794 | + "\n", |
795 | 795 | " if step % 200 == 0:\n", |
796 | 796 | " print(f\"Step {step}/{num_steps} | Total: {loss.item():.2f} | Recon: {recon.item():.2f} | KL: {kl.item():.2f}\")\n", |
797 | 797 | "\n", |
|
0 commit comments