aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorgdamms <damguillotin@gmail.com>2024-04-30 17:15:47 +0200
committergdamms <damguillotin@gmail.com>2024-04-30 17:15:47 +0200
commit021689ca4a825ae2e00ccab2828fed0a34284cba (patch)
treee086c9539c6a5d1497c3f786bfe8d1af5516ea9f
parentf5d919dfddb7a79082141cfe1884bb2678413a6f (diff)
downloaddiffusion-mnist-021689ca4a825ae2e00ccab2828fed0a34284cba.tar.gz
diffusion-mnist-021689ca4a825ae2e00ccab2828fed0a34284cba.zip
num workers
-rw-r--r--main.py132
1 files changed, 66 insertions, 66 deletions
diff --git a/main.py b/main.py
index e86022c..248a991 100644
--- a/main.py
+++ b/main.py
@@ -160,82 +160,82 @@ def loss(y_pred, y_true):
return nn.MSELoss()(y_pred, y_true)
+if __name__ == '__main__':
+ mnist_data = datasets.MNIST(
+ root="./data",
+ train=True,
+ download=True,
+ transform=transforms.ToTensor(),
+ )
+ img, label = mnist_data[0]
+ img = img.to(DEVICE)
+ fig = plt.figure(figsize=(DIFFU_STEPS, 2))
+ for t in range(1, DIFFU_STEPS):
+ xt_1 = q_xt_x0(img, t - 1).sample()
+ x_t = q_xt_xt_1(xt_1, t).sample()
+ ax = fig.add_subplot(2, DIFFU_STEPS, t)
+ ax.imshow(xt_1[0].cpu(), cmap="gray")
+ ax.axis("off")
+ ax = fig.add_subplot(2, DIFFU_STEPS, DIFFU_STEPS + t)
+ ax.imshow(x_t[0].cpu(), cmap="gray")
+ ax.axis("off")
+ fig.tight_layout()
+ fig.savefig("img.tmp.png")
-mnist_data = datasets.MNIST(
- root="./data",
- train=True,
- download=True,
- transform=transforms.ToTensor(),
-)
-img, label = mnist_data[0]
-img = img.to(DEVICE)
-fig = plt.figure(figsize=(DIFFU_STEPS, 2))
-for t in range(1, DIFFU_STEPS):
- xt_1 = q_xt_x0(img, t - 1).sample()
- x_t = q_xt_xt_1(xt_1, t).sample()
- ax = fig.add_subplot(2, DIFFU_STEPS, t)
- ax.imshow(xt_1[0].cpu(), cmap="gray")
- ax.axis("off")
- ax = fig.add_subplot(2, DIFFU_STEPS, DIFFU_STEPS + t)
- ax.imshow(x_t[0].cpu(), cmap="gray")
- ax.axis("off")
-fig.tight_layout()
-fig.savefig("img.tmp.png")
+ ############
+ # Training #
+ ############
-############
-# Training #
-############
+ torch.multiprocessing.set_start_method("spawn")
-torch.multiprocessing.set_start_method("spawn")
+ # Load the model.
+ model = UNet().to(DEVICE)
+ model.load_state_dict(torch.load('model.pth'))
-# Load the model.
-model = UNet().to(DEVICE)
-model.load_state_dict(torch.load('model.pth'))
+ # Define the optimizer.
+ optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)
-# Define the optimizer.
-optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
+ # Define the trainiself.encodevec = nn.Linear(10, 28*28)ng dataset.
+ train_dataset = MNISTDiffusionDataset(train=True)
+ train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True, num_workers=4)
+ trainer = Trainer()
+ criterion = loss
+ epochs = 30
-# Define the training dataset.
-train_dataset = MNISTDiffusionDataset(train=True)
-train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True, num_workers=0)
-trainer = Trainer()
-criterion = loss
-epochs = 10
+ # # Train the model.
+ # trainer.train(model, train_loader, epochs, optimizer, criterion)
-# # Train the model.
-# trainer.train(model, train_loader, epochs, optimizer, criterion)
+ # # Save the model.
+ # torch.save(model.state_dict(), "model.pth")
-# # Save the model.
-# torch.save(model.state_dict(), "model.pth")
+ ##############
+ # Evaluation #
+ ##############
-##############
-# Evaluation #
-##############
+ nb_plots = 5
+ ti_plots = np.linspace(1, DIFFU_STEPS, nb_plots, dtype=int)
+ n_values = [i for i in range(10)]
-nb_plots = 5
-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)))
-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()
-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()
-
- 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).sample()
- 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")
+ 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).sample()
+ 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")