aboutsummaryrefslogtreecommitdiff
path: root/autoencoder.py
diff options
context:
space:
mode:
Diffstat (limited to 'autoencoder.py')
-rw-r--r--autoencoder.py136
1 files changed, 136 insertions, 0 deletions
diff --git a/autoencoder.py b/autoencoder.py
new file mode 100644
index 0000000..f16fe79
--- /dev/null
+++ b/autoencoder.py
@@ -0,0 +1,136 @@
+import torch
+from torch.utils.data import DataLoader, Dataset
+
+from torchvision import datasets, transforms
+
+import matplotlib.pyplot as plt
+
+from trainer import Trainer
+
+
+class PrintLayer(torch.nn.Module):
+ def forward(self, x):
+ print(x.shape)
+ return x
+
+
+class Autoencoder(torch.nn.Module):
+ def __init__(self, input_dim, latent_dim):
+ 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.ReLU(),
+ torch.nn.Conv2d(16, 16, kernel_size=3, padding=1),
+ torch.nn.ReLU(),
+ torch.nn.MaxPool2d(kernel_size=2),
+ torch.nn.Conv2d(16, 32, kernel_size=3, padding=1),
+ torch.nn.ReLU(),
+ torch.nn.Conv2d(32, 32, kernel_size=3, padding=1),
+ torch.nn.ReLU(),
+ torch.nn.MaxPool2d(kernel_size=2),
+ torch.nn.Conv2d(32, 64, kernel_size=3, padding=1),
+ 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.ReLU(),
+ torch.nn.Unflatten(1, latent_dim),
+ )
+ self.decoder = torch.nn.Sequential(
+ torch.nn.Flatten(),
+ torch.nn.Linear(self.latent_size, 64 * 7 * 7),
+ 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),
+ torch.nn.ReLU(),
+ torch.nn.Conv2d(32, 32, kernel_size=3, padding=1),
+ torch.nn.ReLU(),
+ torch.nn.ConvTranspose2d(32, 16, kernel_size=2, stride=2),
+ torch.nn.ReLU(),
+ torch.nn.Conv2d(16, 16, kernel_size=3, padding=1),
+ torch.nn.ReLU(),
+ torch.nn.Conv2d(16, input_dim, kernel_size=3, padding=1),
+ torch.nn.Sigmoid(),
+ )
+
+ def forward(self, x):
+ x = self.encoder(x)
+ x = self.decoder(x)
+ return x
+
+ def encode(self, x):
+ return self.encoder(x)
+
+ def decode(self, x):
+ return self.decoder(x)
+
+
+class AutoencoderDataset(Dataset):
+ def __init__(self, dataset, device='cpu'):
+ self.dataset = dataset
+ self.device = device
+
+ def __len__(self):
+ return len(self.dataset)
+
+ def __getitem__(self, idx):
+ data = self.dataset[idx][0].to(self.device)
+ return data, data
+
+def main():
+ torch.multiprocessing.set_start_method("spawn")
+ device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
+
+ # Load dataset
+ mnist = datasets.MNIST(
+ root='data',
+ train=True,
+ download=True,
+ transform=transforms.ToTensor(),
+ )
+ dataset = AutoencoderDataset(mnist, device=device)
+ dataloader = DataLoader(dataset, batch_size=64, shuffle=True,
+ 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.to(device)
+
+ # Train model
+ trainer = Trainer()
+ lr = 1e-3
+ epochs = 10
+ optimizer = torch.optim.Adam(model.parameters(), lr=lr)
+ criterion = torch.nn.functional.mse_loss
+ trainer.train(model, dataloader, epochs, optimizer, criterion)
+
+ # Save model
+ torch.save(model.state_dict(), 'autoencoder.pth')
+
+ # Visualize results
+ n = 10
+ with torch.no_grad():
+ plt.figure(figsize=(2*n, 4))
+ 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')
+ plt.axis('off')
+ plt.subplot(2, n, i + n + 1)
+ plt.imshow(x_hat.cpu().squeeze().numpy(), cmap='gray')
+ plt.axis('off')
+ plt.tight_layout()
+ plt.savefig('autoencoder.tmp.png')
+
+
+if __name__ == '__main__':
+ main()