diff options
| -rw-r--r-- | main.py | 206 |
1 files changed, 134 insertions, 72 deletions
@@ -267,94 +267,136 @@ 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) + # 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") + # 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") + # 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.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") + # 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) + # # Backward diffusion + # t_plots = np.linspace(1, DIFFU_STEPS, nb_plots, dtype=int) - n_classes = 10 + # 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) + # 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") + # 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.tensor(vec, device=DEVICE, dtype=torch.int64) - vec = torch.nn.functional.one_hot(vec, num_classes=NB_LABEL).to(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.te######## + # # Test Axel # + # ############# - 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) + # # Which image is the closest? + # plt.figure(figsize=(nb_plots, 2)) - x = x * 0.5 + 0.5 - x = x.clamp(0, 1) + # 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) - 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") + # 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") + + + # ###########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 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) + + # x = x * 0.5 + 0.5 + # x = x.clamp(0, 1) + + # 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 # @@ -397,4 +439,24 @@ if __name__ == '__main__': plt.tight_layout() plt.savefig("axel.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) + + 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) + + 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.show()
\ No newline at end of file |
