aboutsummaryrefslogtreecommitdiff
path: root/plots.py
diff options
context:
space:
mode:
authorgdamms <damguillotin@gmail.com>2024-05-17 15:38:32 +0200
committergdamms <damguillotin@gmail.com>2024-05-17 15:38:32 +0200
commit21759fef8826367f8e08bfc1c8118957730cdfdd (patch)
tree1cb4eabf3a66de52743608100229da5ae617c824 /plots.py
parent6a9fdee33106fb048ec4e1cb428da96f3bbd0b39 (diff)
downloaddiffusion-mnist-21759fef8826367f8e08bfc1c8118957730cdfdd.tar.gz
diffusion-mnist-21759fef8826367f8e08bfc1c8118957730cdfdd.zip
more plots
Diffstat (limited to 'plots.py')
-rw-r--r--plots.py84
1 files changed, 62 insertions, 22 deletions
diff --git a/plots.py b/plots.py
index dd4706b..db5f009 100644
--- a/plots.py
+++ b/plots.py
@@ -2,28 +2,19 @@ import matplotlib.pyplot as plt
from matplotlib.gridspec import GridSpec
import numpy as np
from torchvision import datasets
+import torch
+import os
+from rich.progress import track
+from main import UNet, q_xt_xt_1, p_xt_1_xt
-def q_xt_x0(x0, t):
- alpha_bar = ALPHA_BAR[t]
- mean = np.sqrt(alpha_bar) * x0
- std = np.sqrt(1 - alpha_bar)
- eps = np.random.normal(0, 1, x0.shape)
- xt = mean + std * eps
+os.makedirs('plots', exist_ok=True)
+os.makedirs('plots/diffusion', exist_ok=True)
+os.makedirs('plots/diffusion_inverse', exist_ok=True)
- return xt, eps
-
-def q_xt_xt_1(xt_1, t):
- alpha = ALPHA[t]
- mean = np.sqrt(alpha) * xt_1
- std = np.sqrt(1 - alpha)
-
- eps = np.random.normal(0, 1, xt_1.shape)
- xt = mean + std * eps
-
- return xt, eps
+DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
DIFFU_STEPS = 1000
BETA = np.linspace(1e-4, 2e-2, DIFFU_STEPS)
@@ -46,6 +37,7 @@ plt.savefig('plots/alpha_beta.tmp.png')
mnist = datasets.MNIST('data', train=True, download=True)
img = mnist.data[np.random.randint(0, len(mnist))].numpy() / 255
+img = img * 2 - 1
plt.figure()
plt.imshow(img, cmap='gray')
@@ -68,24 +60,72 @@ x_norm = np.linspace(BIN_MIN, BIN_MAX, 100)
y_norm = norm_dist(x_norm, 0, 1) * 28**2 / NB_BINS * (BIN_MAX - BIN_MIN)
fig = plt.figure(figsize=(10, 5))
+fig.suptitle('Diffusion naturelle')
gs = GridSpec(1, 3, figure=fig)
ax1 = fig.add_subplot(gs[0, 0])
ax2 = fig.add_subplot(gs[0, 1:])
-xt = img
-for t in range(DIFFU_STEPS):
+plots_to_save = np.linspace(1, DIFFU_STEPS, 100).astype(int)
+
+xt = torch.tensor(img, device=DEVICE, dtype=torch.float32).unsqueeze(0).unsqueeze(0)
+for t in track(range(1, DIFFU_STEPS+1)):
xt, eps = q_xt_xt_1(xt, t)
+ if t not in plots_to_save:
+ continue
+
+ xt_numpy = xt.cpu().detach().numpy()[0, 0]
+
ax1.clear()
- ax1.imshow(xt, cmap='gray')
- ax1.set_title(f'xt at t={t}')
+ ax1.imshow(xt_numpy, cmap='gray')
+ ax1.set_title(f'xt at t={t:04d}')
ax1.axis('off')
ax2.clear()
- ax2.hist(xt.flatten(), bins=NB_BINS, range=(BIN_MIN, BIN_MAX))
+ ax2.hist(xt_numpy.flatten(), bins=NB_BINS, range=(BIN_MIN, BIN_MAX))
ax2.plot(x_norm, y_norm, color='red', label='N(0, 1)')
+ ax2.legend()
ax2.set_yscale('log')
ax2.set_ylim(y_norm.min(), 1e3)
ax2.set_title(f'xt histogram')
fig.savefig(f'plots/diffusion/{t:04d}.tmp.png')
+os.system('convert --delay 20 --repeat 0 plots/diffusion/*.png plots/diffusion.tmp.gif')
+
+
+
+model = UNet().to(DEVICE)
+model.load_state_dict(torch.load('model.pth'))
+
+fig = plt.figure(figsize=(10, 5))
+fig.suptitle('Diffusion inverse')
+gs = GridSpec(1, 3, figure=fig)
+ax1 = fig.add_subplot(gs[0, 0])
+ax2 = fig.add_subplot(gs[0, 1:])
+
+xt = torch.randn(1, 1, 28, 28, device=DEVICE)
+vec = torch.zeros(1, 10).to(DEVICE)
+for t in track(range(DIFFU_STEPS, 0, -1)):
+ t_tensor = torch.tensor([[t]], device=DEVICE, dtype=torch.float32)
+ xt = p_xt_1_xt(model, xt, t_tensor, vec)
+
+ if t not in plots_to_save:
+ continue
+
+ xt_numpy = xt.cpu().detach().numpy()[0, 0]
+
+ ax1.clear()
+ ax1.imshow(xt_numpy, cmap='gray')
+ ax1.set_title(f'xt at t={t:04d}')
+ ax1.axis('off')
+
+ ax2.clear()
+ ax2.hist(xt_numpy.flatten(), bins=NB_BINS, range=(BIN_MIN, BIN_MAX))
+ ax2.plot(x_norm, y_norm, color='red', label='N(0, 1)')
+ ax2.legend()
+ ax2.set_yscale('log')
+ ax2.set_ylim(y_norm.min(), 1e3)
+ ax2.set_title(f'xt histogram')
+
+ fig.savefig(f'plots/diffusion_inverse/{t:04d}.tmp.png')
+os.system('convert --delay 20 --repeat 0 -reverse plots/diffusion_inverse/*.png plots/diffusion_inverse.tmp.gif')