From 021689ca4a825ae2e00ccab2828fed0a34284cba Mon Sep 17 00:00:00 2001 From: gdamms Date: Tue, 30 Apr 2024 17:15:47 +0200 Subject: num workers --- main.py | 158 ++++++++++++++++++++++++++++++++-------------------------------- 1 file changed, 79 insertions(+), 79 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) - -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 # -############ - -torch.multiprocessing.set_start_method("spawn") - -# 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-3) - -# 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) - -# # Save the model. -# torch.save(model.state_dict(), "model.pth") - -############## -# Evaluation # -############## - -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))) - -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") +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") + + + ############ + # Training # + ############ + + torch.multiprocessing.set_start_method("spawn") + + # 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 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 + + # # Train the model. + # trainer.train(model, train_loader, epochs, optimizer, criterion) + + # # Save the model. + # torch.save(model.state_dict(), "model.pth") + + ############## + # Evaluation # + ############## + + 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))) + + 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") -- cgit v1.3.1