CoolFace
Apppublic

PhenoMENON2025/ScaleNet-AIUpscaler

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
1likes
train_model.py132 linesDownload Raw Back to root
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