aboutsummaryrefslogtreecommitdiff
path: root/autoencoder.py
diff options
context:
space:
mode:
authorgdamms <damguillotin@gmail.com>2024-06-06 10:54:20 +0200
committergdamms <damguillotin@gmail.com>2024-06-06 10:54:20 +0200
commit3b82ce5f658ab27ed6e6eaddf239c553f9431b3e (patch)
tree1c73c942a8f61dd9588cde2f84b8ecf58efe131c /autoencoder.py
parentf9e9312b61608c38e3b4dce86c04df7858aee778 (diff)
downloaddiffusion-mnist-3b82ce5f658ab27ed6e6eaddf239c553f9431b3e.tar.gz
diffusion-mnist-3b82ce5f658ab27ed6e6eaddf239c553f9431b3e.zip
working ldm
Diffstat (limited to 'autoencoder.py')
-rw-r--r--autoencoder.py52
1 files changed, 34 insertions, 18 deletions
diff --git a/autoencoder.py b/autoencoder.py
index 645f628..58326fe 100644
--- a/autoencoder.py
+++ b/autoencoder.py
@@ -21,9 +21,8 @@ class Autoencoder(torch.nn.Module):
super(Autoencoder, self).__init__()
self.input_dim = input_dim
self.latent_dim = latent_dim
- self.latent_size = torch.prod(torch.tensor(latent_dim))
self.encoder = torch.nn.Sequential(
- torch.nn.Conv2d(input_dim[0], 16, kernel_size=3, padding=1),
+ torch.nn.Conv2d(input_dim, 16, kernel_size=3, padding=1),
torch.nn.ReLU(),
torch.nn.Conv2d(16, 16, kernel_size=3, padding=1),
torch.nn.ReLU(),
@@ -37,16 +36,19 @@ class Autoencoder(torch.nn.Module):
torch.nn.ReLU(),
torch.nn.Conv2d(64, 64, kernel_size=3, padding=1),
torch.nn.ReLU(),
- torch.nn.Flatten(),
- torch.nn.Linear(64 * 7 * 7, self.latent_size),
+ torch.nn.Conv2d(64, latent_dim, kernel_size=3, padding=1),
+ torch.nn.ReLU(),
+ # 1x7x7 to 1x8x8
+ torch.nn.Conv2d(latent_dim, latent_dim, kernel_size=2, padding=1),
torch.nn.Sigmoid(),
- torch.nn.Unflatten(1, latent_dim),
)
self.decoder = torch.nn.Sequential(
- torch.nn.Flatten(),
- torch.nn.Linear(self.latent_size, 64 * 7 * 7),
+ # 1x8x8 to 1x7x7
+ torch.nn.Conv2d(latent_dim, latent_dim, kernel_size=2, padding=0),
+ torch.nn.ReLU(),
+ # main decoder
+ torch.nn.Conv2d(latent_dim, 64, kernel_size=3, padding=1),
torch.nn.ReLU(),
- torch.nn.Unflatten(1, (64, 7, 7)),
torch.nn.Conv2d(64, 64, kernel_size=3, padding=1),
torch.nn.ReLU(),
torch.nn.ConvTranspose2d(64, 32, kernel_size=2, stride=2),
@@ -57,7 +59,7 @@ class Autoencoder(torch.nn.Module):
torch.nn.ReLU(),
torch.nn.Conv2d(16, 16, kernel_size=3, padding=1),
torch.nn.ReLU(),
- torch.nn.Conv2d(16, input_dim[0], kernel_size=3, padding=1),
+ torch.nn.Conv2d(16, input_dim, kernel_size=3, padding=1),
torch.nn.Sigmoid(),
)
@@ -102,14 +104,14 @@ def main():
num_workers=4, persistent_workers=True)
# Initialize model
- model = Autoencoder(input_dim=(1, 28, 28), latent_dim=(1, 8, 8))
- model.load_state_dict(torch.load('autoencoder.pth'))
+ model = Autoencoder(input_dim=1, latent_dim=1)
+ model.load_state_dict(torch.load('mnist_autoencoder.pth'))
model.to(device)
# Train model
trainer = Trainer()
lr = 1e-3
- epochs = 10
+ epochs = 1
optimizer = torch.optim.Adam(model.parameters(), lr=lr)
criterion = torch.nn.functional.binary_cross_entropy
trainer.train(model, dataloader, epochs, optimizer, criterion)
@@ -120,17 +122,31 @@ def main():
# Visualize results
n = 10
with torch.no_grad():
- plt.figure(figsize=(2*n, 4))
+ plt.figure(figsize=(2*n, 6))
for i, j in enumerate(torch.randint(0, len(dataset), (n,))):
x, _ = dataset[j]
x = x.unsqueeze(0)
- x_hat = model(x)
- plt.subplot(2, n, i + 1)
- plt.imshow(x.cpu().squeeze().numpy(), cmap='gray')
+ x_latent = model.encode(x)
+ x_hat = model.decode(x_latent)
+ plt.subplot(3, n+1, i + 2)
+ plt.imshow(x.cpu().squeeze().numpy())
+ plt.axis('off')
+ plt.subplot(3, n+1, i + n + 3)
+ plt.imshow(x_latent.cpu().squeeze().numpy())
plt.axis('off')
- plt.subplot(2, n, i + n + 1)
- plt.imshow(x_hat.cpu().squeeze().numpy(), cmap='gray')
+ plt.subplot(3, n+1, i + 2*n + 4)
+ plt.imshow(x_hat.cpu().squeeze().numpy())
plt.axis('off')
+ plt.subplot(3, n+1, 1)
+ plt.text(0.5, 0.5, 'Original', horizontalalignment='center', fontsize=12)
+ plt.axis('off')
+ plt.subplot(3, n+1, n + 2)
+ plt.text(0.5, 0.5, 'Latent', horizontalalignment='center', fontsize=12)
+ plt.axis('off')
+ plt.subplot(3, n+1, 2*n + 3)
+ plt.text(0.5, 0.5, 'Reconstructed', horizontalalignment='center', fontsize=12)
+ plt.axis('off')
+ plt.suptitle('Autoencoder')
plt.tight_layout()
plt.savefig('autoencoder.tmp.png')