PhenoMENON2025/ScaleNet-AIUpscaler
1
1import torch2from torch import nn, optim3from torch.utils.data import DataLoader4from fullsrgan import SRGANGenerator, SRGANDiscriminator5from losses import VGGPerceptualLoss6from dataset_loader import SRDataset7from torch.cuda.amp import autocast, GradScaler8import matplotlib.pyplot as plt9import os10 11torch.backends.cudnn.benchmark = True 12 13 14epochs = 6015batch_size = 216lr = 1e-417lambda_mse = 1.018lambda_adv = 1e-319lambda_perc = 0.00620early_stopping_patience = 1021best_loss = float("inf")22patience_counter = 023 24device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')25scaler = GradScaler() 26 27 28dataset = SRDataset(29 lr_dirs=['data/lr720-1080ver'], 30 hr_dir='data/hr',31 resize_to=(1920, 1080) 32)33loader = DataLoader(34 dataset,35 batch_size=batch_size,36 shuffle=True,37 num_workers=8,38 pin_memory=True,39 persistent_workers=True40)41 42generator = SRGANGenerator(upscale=4, refinement=True).to(device)43discriminator = SRGANDiscriminator().to(device)44 45mse_loss = nn.MSELoss()46adv_loss = nn.BCEWithLogitsLoss() 47vgg_loss = VGGPerceptualLoss().to(device)48 49opt_g = optim.Adam(generator.parameters(), lr=lr)50opt_d = optim.Adam(discriminator.parameters(), lr=lr)51 52log_g_losses, log_d_losses = [], []53 54 55for epoch in range(epochs):56 generator.train()57 discriminator.train()58 total_g_loss, total_d_loss = 0.0, 0.059 60 for lr_img, hr_img in loader:61 lr_img, hr_img = lr_img.to(device), hr_img.to(device)62 63 64 with autocast():65 sr_img = generator(lr_img).detach()66 d_real = discriminator(hr_img)67 d_fake = discriminator(sr_img)68 69 real_labels = torch.ones_like(d_real).to(device)70 fake_labels = torch.zeros_like(d_fake).to(device)71 72 loss_d_real = adv_loss(d_real, real_labels)73 loss_d_fake = adv_loss(d_fake, fake_labels)74 d_loss = (loss_d_real + loss_d_fake) / 275 76 opt_d.zero_grad()77 scaler.scale(d_loss).backward()78 scaler.step(opt_d)79 scaler.update()80 81 82 with autocast():83 sr_img = generator(lr_img)84 d_fake = discriminator(sr_img)85 86 loss_g_adv = adv_loss(d_fake, real_labels)87 loss_g_mse = mse_loss(sr_img, hr_img)88 loss_g_perc = vgg_loss(sr_img, hr_img)89 90 g_loss = lambda_mse * loss_g_mse + lambda_adv * loss_g_adv + lambda_perc * loss_g_perc91 92 opt_g.zero_grad()93 scaler.scale(g_loss).backward()94 scaler.step(opt_g)95 scaler.update()96 97 total_d_loss += d_loss.item()98 total_g_loss += g_loss.item()99 log_g_losses.append(g_loss.item())100 log_d_losses.append(d_loss.item())101 102 103 avg_g_loss = total_g_loss / len(loader)104 avg_d_loss = total_d_loss / len(loader)105 print(f"[Epoch {epoch+1}/{epochs}] G Loss: {avg_g_loss:.4f} | D Loss: {avg_d_loss:.4f}")106 107 108 if avg_g_loss < best_loss:109 best_loss = avg_g_loss110 patience_counter = 0111 torch.save(generator.state_dict(), "best_srgan_generator4.pth")112 else:113 patience_counter += 1114 if patience_counter >= early_stopping_patience:115 print("Early stopping triggered")116 break117 118 119torch.save(generator.state_dict(), "srgan_generator4.pth")120 121 122plt.figure(figsize=(10, 5))123plt.plot(log_g_losses, label="Generator Loss")124plt.plot(log_d_losses, label="Discriminator Loss")125plt.xlabel("Training Iterations")126plt.ylabel("Loss")127plt.title("SRGAN Loss Curve")128plt.legend()129plt.grid(True)130plt.savefig("loss_plot.png")131plt.show()132 