aboutsummaryrefslogtreecommitdiff
path: root/main.py
diff options
context:
space:
mode:
authorgdamms <damguillotin@gmail.com>2024-07-02 11:39:08 +0200
committergdamms <damguillotin@gmail.com>2024-07-02 11:39:08 +0200
commit9d6d3f3c7f8f01e24d635c3f5d9b43fa697d5f9f (patch)
tree803c7a391acf67b2b4f37006a123af2f583f8ff3 /main.py
parent965b433ac9e0b18c22cda8df1058876fe08dfb19 (diff)
downloaddiffusion-mnist-9d6d3f3c7f8f01e24d635c3f5d9b43fa697d5f9f.tar.gz
diffusion-mnist-9d6d3f3c7f8f01e24d635c3f5d9b43fa697d5f9f.zip
unteste update using troch-trainer
Diffstat (limited to 'main.py')
-rw-r--r--main.py47
1 files changed, 27 insertions, 20 deletions
diff --git a/main.py b/main.py
index 7473df1..50c3930 100644
--- a/main.py
+++ b/main.py
@@ -7,13 +7,15 @@ from rich.progress import track
from torchvision import datasets, transforms
+from trainer import train
+from trainer.trainer import Trainer
+
import matplotlib.pyplot as plt
import numpy as np
import os
import cv2
-from trainer import Trainer
from autoencoder import Autoencoder
from utils import *
@@ -33,7 +35,7 @@ class SelfAttention(nn.Module):
return x
-class UNet(nn.Module):
+class UNetEDF(nn.Module):
def __init__(self):
super().__init__()
@@ -262,12 +264,19 @@ if autoencoder is not None:
NB_CHANNEL, IMG_SIZE, _ = img.shape
NB_LABEL = 1
-EPOCHS = 200
+EPOCHS = 1
LEARNING_RATE = 2e-4
-def epoch_callback(epoch_i, epochs, model, trainer):
- if epoch_i % 10 == 0 or epoch_i == epochs - 1:
+def epoch_callback(trainer: Trainer):
+ epoch_i = trainer.epoch_i
+ epochs = trainer.epochs
+
+ if epoch_i % 10 == 0 or epoch_i == epochs:
+ save_path = f'runs/{trainer.run_name}/checkpoints/{epoch_i:04}e.pt'
+ torch.save(trainer.model, save_path)
+ print(f"Model saved at {save_path}")
+
print("Calculating metrics...")
with torch.no_grad():
batch_size = 64
@@ -289,10 +298,12 @@ def epoch_callback(epoch_i, epochs, model, trainer):
reals = torch.stack([dataset[i][0] for i in range(n_samples)]).cpu().numpy()
reals = reals * 2 - 1
- trainer.writer.add_scalars('Metrics/FID', {trainer.date_time: fid(reals, fakes)}, epoch_i)
- trainer.writer.add_scalars('Metrics/KL', {trainer.date_time: kl(reals, fakes)}, epoch_i)
- trainer.writer.add_scalars('Metrics/RKL', {trainer.date_time: kl(fakes, reals)}, epoch_i)
- trainer.writer.add_scalars('Metrics/JSD', {trainer.date_time: jsd(reals, fakes)}, epoch_i)
+ trainer.writer.add_scalar('FID/Validation', fid(reals, fakes), epoch_i)
+ trainer.writer.add_scalar('KL/Validation', kl(reals, fakes), epoch_i)
+ trainer.writer.add_scalar('RKL/Validation', kl(fakes, reals), epoch_i)
+ trainer.writer.add_scalar('JSD/Validation', jsd(reals, fakes), epoch_i)
+
+ trainer.writer.add_images('Fakes/Validation', fakes[:16], epoch_i)
if __name__ == '__main__':
@@ -304,12 +315,12 @@ if __name__ == '__main__':
torch.multiprocessing.set_start_method("spawn")
# Load the model.
- model = UNet().to(DEVICE)
- try:
- model.load_state_dict(torch.load('model.pth'))
- except FileNotFoundError:
- print("No model found, training a new one.")
- pass
+ model = UNetEDF().to(DEVICE)
+ # try:
+ # model = torch.load('model.pth')
+ # except FileNotFoundError:
+ # print("No model found, training a new one.")
+ # pass
# Define the optimizer.
optimizer = torch.optim.Adam(model.parameters(), lr=LEARNING_RATE)
@@ -318,15 +329,11 @@ if __name__ == '__main__':
train_dataset = DiffusionDataset(dataset, autoencoder)
train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True,
num_workers=4, persistent_workers=True)
- trainer = Trainer()
criterion = loss
epochs = EPOCHS
# Train the model.
- trainer.train(model, train_loader, epochs, optimizer, criterion, epoch_callbacks=[epoch_callback])
-
- # Save the model.
- torch.save(model.state_dict(), 'model.pth')
+ train(model, train_loader, epochs, optimizer, criterion, epoch_callbacks=[epoch_callback], save_chekpoint=False)
##############
# Evaluation #