aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--main.py235
1 files changed, 66 insertions, 169 deletions
diff --git a/main.py b/main.py
index 4bbf7c5..2d5c9f1 100644
--- a/main.py
+++ b/main.py
@@ -267,196 +267,93 @@ if __name__ == '__main__':
with torch.no_grad():
- # # Forward diffusion
- # img, label = dataset[np.random.randint(0, len(dataset))]
- # img = img.to(DEVICE) * 2 - 1
+ # Forward diffusion
+ img, label = dataset[np.random.randint(0, len(dataset))]
+ img = img.to(DEVICE) * 2 - 1
- # nb_plots = 10
- # plots_id = [i for i in np.linspace(1, DIFFU_STEPS, nb_plots, dtype=int)]
+ nb_plots = 10
+ 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.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 + 1, plot_i + 2)
- # plt.title(f"t={t}")
- # plt.imshow(tensor_to_image(x))
- # plt.axis("off")
-
- # for t in range(1, DIFFU_STEPS+1):
- # x = q_xt_x0(img, t)[0].cpu()
- # if t not in plots_id:
- # continue
- # plot_i = plots_id.index(t)
- # plt.subplot(2, nb_plots + 1, nb_plots + plot_i + 3)
- # plt.imshow(tensor_to_image(x))
- # 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")
-
-
- # # Backward diffusion
- # t_plots = np.linspace(1, DIFFU_STEPS, nb_plots, dtype=int)
-
- # n_classes = 10
-
- # x = torch.randn(n_classes, NB_CHANNEL, IMG_SIZE, IMG_SIZE, device=DEVICE)
- # vec = torch.tensor([[min(i, NB_LABEL-1)] for i in range(n_classes)], dtype=torch.int64)
- # vec = torch.nn.functional.one_hot(vec, num_classes=NB_LABEL).to(device=DEVICE, dtype=torch.float32)
-
- # 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(tensor_to_image(x[class_i]))
- # plt.axis("off")
- # plt.tight_layout()
- # plt.savefig("backward_diffusion.tmp.png")
-
-
- # # Benchmark
- # x = torch.randn(nb_plots * n_classes, NB_CHANNEL, IMG_SIZE, IMG_SIZE).to(DEVICE)
- # vec = sum([[[min(i, NB_LABEL-1)]] * nb_plots for i in range(n_classes)], [])
- # vec = torch.te########
- # # Test Axel #
- # #############
-
- # # Which image is the closest?
- # plt.figure(figsize=(nb_plots, 2))
-
- # x = torch.randn(nb_plots, NB_CHANNEL, IMG_SIZE, IMG_SIZE).to(DEVICE)
- # vec = torch.randint(0, NB_LABEL, (nb_plots,), device=DEVICE)
- # vec = torch.nn.functional.one_hot(vec, num_classes=NB_LABEL).to(device=DEVICE, dtype=torch.float32)
-
- # for ti in range(DIFFU_STEPS, 0, -1):
- # t = torch.tensor([[ti]] * nb_plots, device=DEVICE, dtype=torch.float32)
- # x = p_xt_1_xt(model, x, t, vec)
-
- # x = x * 0.5 + 0.5
- # x = x.clamp(0, 1)
-
- # for plot_i in range(nb_plots):
- # xi = x[plot_i]
- # imgs = dataset.data.to(DEVICE).to(dtype=torch.float32) / 255
- # dist = torch.norm(imgs - xi, dim=(1, 2))
- # closest_i = torch.argmin(dist)
- # closest_img = imgs[closest_i].unsqueeze(0)
-
- # plt.subplot(2, nb_plots + 1, 2 + plot_i)
- # plt.imshow(tensor_to_image(xi))
- # plt.axis("off")
- # plt.subplot(2, nb_plots + 1 , nb_plots + 3 + plot_i)
- # plt.imshow(tensor_to_image(closest_img))
- # plt.axis("off")
- # plt.suptitle("Closest image")
- # plt.subplot(2, nb_plots + 1, 1)
- # plt.text(0, 0.5, "Generated", fontsize=12)
- # plt.axis("off")
- # plt.subplot(2, nb_plots + 1, nb_plots + 2)
- # plt.text(0, 0.5, "Closest", fontsize=12)
- # plt.axis("off")
- # plt.tight_layout()
- # plt.savefig("axel.tmp.png")
+ xs = forward_diffusion(img)
+ 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 + 1, plot_i + 2)
+ plt.title(f"t={t}")
+ plt.imshow(tensor_to_image(x))
+ plt.axis("off")
- # ###########nsor(vec, device=DEVICE, dtype=torch.int64)
- # vec = torch.nn.functional.one_hot(vec, num_classes=NB_LABEL).to(device=DEVICE, dtype=torch.float32)
+ for t in range(1, DIFFU_STEPS+1):
+ x = q_xt_x0(img, t)[0].cpu()
+ if t not in plots_id:
+ continue
+ plot_i = plots_id.index(t)
+ plt.subplot(2, nb_plots + 1, nb_plots + plot_i + 3)
+ plt.imshow(tensor_to_image(x))
+ plt.axis("off")
- # for ti in range(DIFFU_STEPS, 0, -1):
- # t = torch.tensor([[ti]] * n_classes * nb_plots, device=DEVICE, dtype=torch.float32)
- # x = p_xt_1_xt(model, x, t, vec)
+ 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")
- # x = x * 0.5 + 0.5
- # x = x.clamp(0, 1)
+ plt.suptitle("Forward diffusion")
+ plt.tight_layout()
+ plt.savefig("forward_diffusion.tmp.png")
- # plt.figure(figsize=(nb_plots, n_classes))
- # for i in range(nb_plots):
- # for j in range(n_classes):
- # id = i * n_classes + j
- # plt.subplot(n_classes, nb_plots, id + 1)
- # plt.imshow(tensor_to_image(x[id]))
- # plt.axis("off")
- # plt.tight_layout()
- # plt.savefig("benchmark.tmp.png")
- #############
- # Test Axel #
- #############
+ # Backward diffusion
+ t_plots = np.linspace(1, DIFFU_STEPS, nb_plots, dtype=int)
- # Which image is the closest?
- plt.figure(figsize=(nb_plots, 2))
+ n_classes = 10
- x = torch.randn(nb_plots, NB_CHANNEL, IMG_SIZE, IMG_SIZE).to(DEVICE)
- vec = torch.randint(0, NB_LABEL, (nb_plots,), device=DEVICE)
+ x = torch.randn(n_classes, NB_CHANNEL, IMG_SIZE, IMG_SIZE, device=DEVICE)
+ vec = torch.tensor([[min(i, NB_LABEL-1)] for i in range(n_classes)], dtype=torch.int64)
vec = torch.nn.functional.one_hot(vec, num_classes=NB_LABEL).to(device=DEVICE, dtype=torch.float32)
- for ti in range(DIFFU_STEPS, 0, -1):
- t = torch.tensor([[ti]] * nb_plots, device=DEVICE, dtype=torch.float32)
- x = p_xt_1_xt(model, x, t, vec)
-
- x = x * 0.5 + 0.5
- x = x.clamp(0, 1)
-
- for plot_i in range(nb_plots):
- xi = x[plot_i]
- imgs = dataset.data.to(DEVICE).to(dtype=torch.float32) / 255
- dist = torch.norm(imgs - xi, dim=(1, 2))
- closest_i = torch.argmin(dist)
- closest_img = imgs[closest_i].unsqueeze(0)
-
- plt.subplot(2, nb_plots + 1, 2 + plot_i)
- plt.imshow(tensor_to_image(xi))
- plt.axis("off")
- plt.subplot(2, nb_plots + 1 , nb_plots + 3 + plot_i)
- plt.imshow(tensor_to_image(closest_img))
- plt.axis("off")
- plt.suptitle("Closest image")
- plt.subplot(2, nb_plots + 1, 1)
- plt.text(0, 0.5, "Generated", fontsize=12)
- plt.axis("off")
- plt.subplot(2, nb_plots + 1, nb_plots + 2)
- plt.text(0, 0.5, "Closest", fontsize=12)
- plt.axis("off")
+ 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(tensor_to_image(x[class_i]))
+ plt.axis("off")
plt.tight_layout()
- plt.savefig("axel.tmp.png")
-
+ plt.savefig("backward_diffusion.tmp.png")
- ###############
- # Test Thomas #
- ###############
- x = torch.randn(1, NB_CHANNEL, IMG_SIZE, IMG_SIZE).to(DEVICE)
- vec = torch.tensor([[0, 0, 0, 0, 0, 0, 0, 0, 0, 0]], device=DEVICE, dtype=torch.float32)
+ # Benchmark
+ x = torch.randn(nb_plots * n_classes, NB_CHANNEL, IMG_SIZE, IMG_SIZE).to(DEVICE)
+ vec = sum([[[min(i, NB_LABEL-1)]] * 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=NB_LABEL).to(device=DEVICE, dtype=torch.float32)
for ti in range(DIFFU_STEPS, 0, -1):
- t = torch.tensor([[ti]], 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)
x = x * 0.5 + 0.5
x = x.clamp(0, 1)
- plt.figure(figsize=(1, 1))
- plt.imshow(tensor_to_image(x[0]))
- plt.axis("off")
- plt.savefig("thomas.tmp.png")
+ plt.figure(figsize=(nb_plots, n_classes))
+ for i in range(nb_plots):
+ for j in range(n_classes):
+ id = i * n_classes + j
+ plt.subplot(n_classes, nb_plots, id + 1)
+ plt.imshow(tensor_to_image(x[id]))
+ plt.axis("off")
+ plt.tight_layout()
+ plt.savefig("benchmark.tmp.png")
# plt.show() \ No newline at end of file