### 1. Dependencies

In [1]:
import torch
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms
from torch.autograd import Variable
import torch.optim as optim
import torch.nn as nn
import torch.nn.functional as F
from torch.nn import DataParallel

import time
import os
import numpy as np
import json
import cv2
from PIL import Image, ImageOps
import random
from tqdm import tqdm
import operator
import itertools
from scipy.io import  loadmat
import logging
from scipy import signal

from utils import data_transforms
from utils import get_paste_kernel, kernel_map
from utils_logging import setup_logger

### 2. Choose between Recasens or GazeNet

- Idea is you can just swap 
models.recasens, dataloader.recasens, training.train_recasens, etc...
- with the following
models.gazenet, dataloader.gazenet, training.train_gazenet

In [2]:
from models.gazenet import GazeNet
from models.__init__ import save_checkpoint, resume_checkpoint
from dataloader.gazenet import GooDataset, GazeDataset
from training.train_gazenet import train, test, GazeOptimizer

In [3]:
# Logger will save the training and test errors to a .log file 
logger = setup_logger(name='first_logger', 
                      log_dir ='./logs/',
                      log_file='train_gazenet.log',
                      log_format = '%(asctime)s %(levelname)s %(message)s',
                      verbose=True)

### 3. Dataloaders
- Choose between GazeDataset (Gazefollow dataset) or GooDataset (GooSynth/GooReal)
- Set paths to image directories and pickle paths. For Gazefollow, images_dir and test_images_dir should be the same and both lead to the path containing the train and test folders.

In [4]:
# Dataloaders for GOO
batch_size=32
workers=12
testbatchsize=32

images_dir = '/hdd/HENRI/goosynth/1person/GazeDatasets/'
pickle_path = '/hdd/HENRI/goosynth/picklefiles/trainpickle2to19human.pickle'
test_images_dir = '/hdd/HENRI/goosynth/test/'
test_pickle_path = '/hdd/HENRI/goosynth/picklefiles/testpickle120.pickle'

train_set = GooDataset(images_dir, pickle_path, 'train', use_gazemask=True)
train_data_loader = torch.utils.data.DataLoader(train_set, batch_size=batch_size, shuffle=True, num_workers=workers)

val_set = GooDataset(test_images_dir, test_pickle_path, 'test')
test_data_loader = torch.utils.data.DataLoader(val_set, batch_size=testbatchsize, num_workers=workers, shuffle=False)

Number of Images: 172800
Number of Images: 19200


### 4. Load Model and Set Training Hyperparameters
- For Gazefollow, the model requires the alexnet_places365 pretrained model, provided here: https://urlzs.com/ytKK3
- When resuming training, set to True and set the resume_path for the saved model.
- Here, logging module is initialized (logger) to save training and testing errors.

In [5]:
# Loads model
net = GazeNet()
net.cuda()

# Hyperparameters
start_epoch = 0
max_epoch = 25
learning_rate = 0.0001

# Initializes Optimizer
gaze_opt = GazeOptimizer(net, learning_rate)
optimizer = gaze_opt.getOptimizer(start_epoch)

# Is training resumed? If so, set the resume_path and set flag to True
# This can also be used to evaluate a model 
resume_training = False
resume_path = './saved_models/gazenet_goo/model_epoch25.pth.tar'
if resume_training :
    net, optimizer, start_epoch = resume_checkpoint(net, optimizer, resume_path)
    test(net, test_data_loader,logger)

### 5. Training the Model
- Determine in which epochs do you want to save the model, as you might not want to save every epoch
- Training and test errors can be accessed in the logs directory set up earlier

In [7]:
for epoch in range(start_epoch, max_epoch):
    
    # Update optimizer
    optimizer = gaze_opt.getOptimizer(epoch)

    # Train model
    train(net, train_data_loader, optimizer, epoch, logger)

    # Save model and optimizer
    if epoch > max_epoch-5:
        save_path = './saved_models/gazemask/'
        save_checkpoint(net, optimizer, epoch+1, save_path)
    
    # Evaluate model
    test(net, test_data_loader, logger)

  2%|▏         | 99/5400 [02:33<1:16:41,  1.15it/s][0.66807419 0.58496901 0.58496901]
  4%|▎         | 199/5400 [04:48<1:33:21,  1.08s/it][0.66792659 0.20656406 0.20656406]
  6%|▌         | 299/5400 [06:49<39:41,  2.14it/s]  [0.66798574 0.12672718 0.12672718]
  7%|▋         | 399/5400 [08:55<4:24:01,  3.17s/it][0.6678073  0.12175475 0.12175475]
  9%|▉         | 499/5400 [10:52<1:20:03,  1.02it/s][0.66778867 0.07701725 0.07701725]
 11%|█         | 599/5400 [12:43<40:53,  1.96it/s]  [0.66792505 0.09797143 0.09797143]
 13%|█▎        | 699/5400 [14:46<3:45:51,  2.88s/it][0.66785793 0.08412297 0.08412297]
 15%|█▍        | 799/5400 [16:45<3:25:57,  2.69s/it][0.66795868 0.07445891 0.07445891]
 17%|█▋        | 899/5400 [18:35<1:20:42,  1.08s/it][0.6679272  0.07688928 0.07688928]
 18%|█▊        | 999/5400 [20:28<3:38:23,  2.98s/it][0.6680243  0.06324008 0.06324008]
 20%|██        | 1099/5400 [22:21<3:13:35,  2.70s/it][0.66796541 0.0615291  0.0615291 ]
 22%|██▏       | 1199/5400 [24:12<59:02,  1

KeyboardInterrupt: 