aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--main.py155
1 files changed, 84 insertions, 71 deletions
diff --git a/main.py b/main.py
index 60d453f..e5fc5dd 100644
--- a/main.py
+++ b/main.py
@@ -186,24 +186,56 @@ if __name__ == '__main__':
download=True,
transform=transforms.ToTensor(),
)
+ ############
+ # Training #
+ ############
+
+ torch.multiprocessing.set_start_method("spawn")
+
+ # Load the model.
+ model = UNet().to(DEVICE)
+ try:
+ model.load_state_dict(torch.load('model.pth'))
+ except FileNotFoundError:
+ print("No model found, training a new one.")
+ pass
+
+ # Define the optimizer.
+ optimizer = torch.optim.Adam(model.parameters(), lr=2e-4)
+
+ # Define the training dataset.
+ train_dataset = MNISTDiffusionDataset(train=True)
+ train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True,
+ num_workers=4, persistent_workers=True)
+ trainer = Trainer()
+ criterion = loss
+ epochs = 0
+
+ # Train the model.
+ trainer.train(model, train_loader, epochs, optimizer, criterion)
+
+ # Save the model.
+ torch.save(model.state_dict(), 'model.pth')
- #############
- # Diffusion #
- #############
+ ##############
+ # Evaluation #
+ ##############
+
+ # Forward diffusion
img, label = mnist_data[np.random.randint(0, len(mnist_data))]
img = img.to(DEVICE) * 2 - 1
- nb_plots = 10
+ nb_plots = 6
plots_id = [i for i in np.linspace(1, DIFFU_STEPS, nb_plots, dtype=int)]
xs = forward_diffusion(img)
- plt.figure(figsize=(nb_plots, 2))
+ plt.figure(figsize=(nb_plots, 2.5))
for t, x in enumerate(xs):
if t not in plots_id:
continue
plot_i = plots_id.index(t)
- plt.subplot(2, nb_plots, plot_i + 1)
+ plt.subplot(2, nb_plots + 1, plot_i + 2)
plt.title(f"t={t}")
plt.imshow(x, cmap="gray")
plt.axis("off")
@@ -213,86 +245,67 @@ if __name__ == '__main__':
if t not in plots_id:
continue
plot_i = plots_id.index(t)
- plt.subplot(2, nb_plots, nb_plots + plot_i + 1)
+ plt.subplot(2, nb_plots + 1, nb_plots + plot_i + 3)
plt.imshow(x, cmap="gray")
plt.axis("off")
+ plt.subplot(2, nb_plots + 1, 1)
+ plt.text(0, 0.5, "Implicit", fontsize=12)
+ plt.axis("off")
+ plt.subplot(2, nb_plots + 1, nb_plots + 2)
+ plt.text(0, 0.5, "Explicit", fontsize=12)
+ plt.axis("off")
+
plt.suptitle("Forward diffusion")
plt.tight_layout()
plt.savefig("forward_diffusion.tmp.png")
- exit()
- ############
- # Training #
- ############
- torch.multiprocessing.set_start_method("spawn")
+ # Backward diffusion
+ nb_plots = 6
+ t_plots = np.linspace(1, DIFFU_STEPS, nb_plots, dtype=int)
- # Load the model.
- model = UNet().to(DEVICE)
- model.load_state_dict(torch.load('model.pth'))
+ n_classes = 10
- # Define the optimizer.
- optimizer = torch.optim.Adam(model.parameters(), lr=2e-4)
+ x = torch.randn(n_classes, 1, 28, 28,device=DEVICE)
+ vec = torch.tensor([[i] for i in range(n_classes)], dtype=torch.int64)
+ vec = torch.nn.functional.one_hot(vec, num_classes=10).to(device=DEVICE, dtype=torch.float32)
- # Define the training dataset.
- train_dataset = MNISTDiffusionDataset(train=True)
- train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True,
- num_workers=4, persistent_workers=True)
- trainer = Trainer()
- criterion = loss
- epochs = 1
-
- # # Train the model.
- # trainer.train(model, train_loader, epochs, optimizer, criterion)
-
- # # Save the model.
- # torch.save(model.state_dict(), 'model.pth')
-
- ##############
- # Evaluation #
- ##############
-
- nb_plots = 10
- ti_plots = np.linspace(1, DIFFU_STEPS, nb_plots, dtype=int)
- n_values = [i for i in range(10)]
-
- fig = plt.figure(figsize=(nb_plots, len(n_values)))
-
- for attempti, n in enumerate(n_values):
- n = torch.tensor([[n]], device=DEVICE, dtype=torch.int64)
- x = torch.randn(1, 1, 28, 28).to(DEVICE)
- vec = torch.nn.functional.one_hot(n, num_classes=10).to(device=DEVICE, dtype=torch.float32)
- img_vec = model.encodevec(vec)
- img_vec = F.relu(img_vec)
- img_vec = img_vec.view(-1, 1, 28, 28)
- img_vec = img_vec.cpu().detach().numpy()
+ plt.figure(figsize=(nb_plots, n_classes))
+ plt.suptitle("Backward diffusion")
+ for t in range(DIFFU_STEPS, 0, -1):
+ t_tensor = torch.tensor([[t]] * n_classes, device=DEVICE, dtype=torch.float32)
+ x = p_xt_1_xt(model, x, t_tensor, vec)
+ if t in t_plots:
+ t_plot_i = nb_plots - t_plots.tolist().index(t) - 1
+ for class_i in range(n_classes):
+ plt.subplot(n_classes, nb_plots, t_plot_i + nb_plots * class_i + 1)
+ if class_i == 0:
+ plt.title(f"t={t}")
+ plt.imshow(x[class_i, 0].detach().cpu(), cmap="gray")
+ plt.axis("off")
+ plt.tight_layout()
+ plt.savefig("backward_diffusion.tmp.png")
- for ti in range(DIFFU_STEPS, 0, -1):
- t = torch.tensor([[ti]], device=DEVICE, dtype=torch.float32)
- x = p_xt_1_xt(model, x, t, vec)
- if ti in ti_plots:
- ti_plotind = nb_plots - np.where(ti_plots == ti)[0][0]
- ax = fig.add_subplot(len(n_values), nb_plots, ti_plotind + nb_plots * attempti)
- ax.imshow(x[0, 0].detach().cpu(), cmap="gray")
- ax.axis("off")
- ax.set_title(f"{ti}")
- fig.tight_layout()
- fig.savefig("diffused.tmp.png")
- x = torch.randn(nb_plots * 10, 1, 28, 28).to(DEVICE)
- values = sum([[[i]] * nb_plots for i in range(10)], [])
- vec = torch.nn.functional.one_hot(torch.tensor(values, device=DEVICE, dtype=torch.int64), num_classes=10).to(device=DEVICE, dtype=torch.float32)
+ # Benchmark
+ x = torch.randn(nb_plots * n_classes, 1, 28, 28).to(DEVICE)
+ vec = sum([[[i]] * nb_plots for i in range(n_classes)], [])
+ vec = torch.tensor(vec, device=DEVICE, dtype=torch.int64)
+ vec = torch.nn.functional.one_hot(vec, num_classes=10).to(device=DEVICE, dtype=torch.float32)
for ti in range(DIFFU_STEPS, 0, -1):
- t = torch.tensor([[ti]] * 10 * nb_plots, device=DEVICE, dtype=torch.float32)
+ t = torch.tensor([[ti]] * n_classes * nb_plots, device=DEVICE, dtype=torch.float32)
x = p_xt_1_xt(model, x, t, vec)
- fig = plt.figure(figsize=(nb_plots, len(n_values)))
+ plt.figure(figsize=(nb_plots, n_classes))
for i in range(nb_plots):
- for j in range(10):
- ax = fig.add_subplot(10, nb_plots, i + nb_plots * j + 1)
- ax.imshow(x[0, 0].detach().cpu(), cmap="gray")
- ax.axis("off")
- fig.tight_layout()
- fig.savefig("diffused_all.tmp.png") \ No newline at end of file
+ for j in range(n_classes):
+ id = i * n_classes + j
+ plt.subplot(n_classes, nb_plots, id + 1)
+ plt.imshow(x[id, 0].detach().cpu(), cmap="gray")
+ plt.axis("off")
+ plt.tight_layout()
+ plt.savefig("benchmark.tmp.png")
+
+ plt.show() \ No newline at end of file