StyleCast / train.py
tejas-m's picture
Initial commit: StyleCast - Neural Style Transfer with AdaIN
c6bf2f6
Raw
History Blame Contribute Delete
7.14 kB
import argparse
import torch
from torch.utils.data import DataLoader
import torch.optim as optim
from pathlib import Path
from utils.utils import *
from utils.models import *
from tqdm import tqdm
from torchvision.utils import save_image
def parse_arguments():
parser = argparse.ArgumentParser()
parser.add_argument('--content_dir', type=str, default=r'C:\Users\Tejas\OneDrive\Desktop\AI Major Projects\Neural Style Transfer with AdaIN\content_data',
help='Location of content dataset')
parser.add_argument('--style_dir', type=str, default=r'C:\Users\Tejas\OneDrive\Desktop\AI Major Projects\Neural Style Transfer with AdaIN\style_data',
help='Location of style dataset')
parser.add_argument('--vgg', type=str, default=r'C:\Users\Tejas\OneDrive\Desktop\AI Major Projects\Neural Style Transfer with AdaIN\vgg_normalised.pth',
help='Location of pre-trained VGG')
parser.add_argument('--experiment', type=str, default='experiment1',
help='Name of experiment')
parser.add_argument('--final_size', type=int, default=256,
help='Size of final image')
parser.add_argument('--content_size', type=int, default=512,
help='Size of content image')
parser.add_argument('--style_size', type=int, default=512,
help='Size of style image')
parser.add_argument('--crop', action='store_true', default=True,
help='Crop image')
parser.add_argument('--batch_size', type=int, default=4,
help='Batch size')
parser.add_argument('--lr', type=float, default=1e-4,
help='Learning rate')
parser.add_argument('--lr_decay', type=float, default=5e-5,
help='Learning rate decay')
parser.add_argument('--epochs', type=int, default=1,
help='Number of epochs')
parser.add_argument('--content_weight', type=float, default=1.0,
help='Content weight')
parser.add_argument('--style_weight', type=float, default=5,
help='Style weight')
parser.add_argument('--log_interval', type=int, default=1,
help='Log interval')
parser.add_argument('--save_interval', type=int, default=2,
help='Save interval')
parser.add_argument('--resume', action='store_true', default=False,
help='Resume training')
parser.add_argument('--decoder_path', type=str, default=None,
help='Path to decoder checkpoint')
parser.add_argument('--optimizer_path', type=str, default=None,
help='Path to optimizer checkpoint')
return parser.parse_args()
def main():
args = parse_arguments()
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
save_dir = Path('experiment') / args.experiment
save_dir.mkdir(exist_ok=True, parents=True)
#Save argument values
with open(save_dir / 'args.txt', 'w') as args_file:
for key, value in vars(args).items():
args_file.write(f'{key}: {value}\n')
content_transform = get_transform(args.content_size, args.crop, args.final_size)
style_transform = get_transform(args.style_size, args.crop, args.final_size)
content_dataset = ImageFolderDataset(args.content_dir, content_transform)
style_dateset = ImageFolderDataset(args.style_dir, style_transform)
content_dataloader = DataLoader(content_dataset,
batch_size=args.batch_size,
shuffle = True,
pin_memory=True,
drop_last=True)
style_dataloader = DataLoader(style_dateset,
batch_size=args.batch_size,
shuffle=True,
pin_memory=True,
drop_last=True)
print('Number of batches in content dataset: ', len(content_dataloader))
print('Number of batches in style dataset: ', len(style_dataloader))
encoder = VGGEncoder(args.vgg).to(device)
decoder = Decoder().to(device)
optimizer = optim.Adam(decoder.parameters(), lr=args.lr)
scheduler = optim.lr_scheduler.LambdaLR(
optimizer,
lr_lambda = lambda epoch: 1.0 / (1.0 + args.lr_decay * epoch)
)
if args.resume:
decoder.load_state_dict(torch.load(args.decoder_path))
optimizer.load_state_dict(torch.load(args.optimizer_path))
print('Training...')
mse_loss = torch.nn.MSELoss()
encoder.eval()
running_loss = None
running_closs = None
running_sloss = None
for epoch in range(args.epochs):
progress_bar = tqdm(zip(content_dataloader, style_dataloader),
total=min(len(content_dataloader), len(style_dataloader)))
running_loss = 0
running_closs = 0
running_sloss = 0
for content_batch, style_batch in progress_bar:
content_batch = content_batch.to(device)
style_batch = style_batch.to(device)
c_feats = encoder(content_batch)
s_feats = encoder(style_batch)
t = adaptive_instance_normalization(c_feats[-1], s_feats[-1])
g = decoder(t)
g_feats = encoder(g)
loss_c = mse_loss(g_feats[-1], t) * args.content_weight
loss_s = 0
for g_f, s_f in zip(g_feats, s_feats):
g_mean, g_std = calc_mean_std(g_f)
s_mean, s_std = calc_mean_std(s_f)
loss_s += mse_loss(g_mean, s_mean) + mse_loss(g_std, s_std)
loss_s = loss_s * args.style_weight
loss = loss_c + loss_s
optimizer.zero_grad()
loss.backward()
optimizer.step()
progress_bar.set_description(f'Loss:{loss.item():4f}, Content Loss: {loss_c.item():4f}, Style Loss: {loss_s.item():4f}')
running_loss += loss.item()
running_closs += loss_c.item()
running_sloss += loss_s.item()
scheduler.step()
running_loss /= len(content_dataloader)
running_closs /= len(content_dataloader)
running_sloss /= len(content_dataloader)
if (epoch+1) % args.log_interval == 0:
tqdm.write(f'Iter {epoch+1}: Loss:{running_loss:4f}, Content Loss: {running_closs:4f}, Style Loss: {running_sloss:4f}')
if (epoch+1) % args.save_interval == 0:
torch.save(decoder.state_dict(), save_dir / f'decoder_{epoch+1}.pth')
torch.save(optimizer.state_dict(), save_dir / f'optimizer_{epoch+1}.pth')
with torch.no_grad():
output = torch.cat([content_batch, style_batch, g], dim=0)
save_image(output, save_dir / f'output_{epoch+1}.png', nrow=args.batch_size)
if __name__ == '__main__':
main()