Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 7 additions & 1 deletion helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ def load_run(run_name):
mnist, cifar = pd.read_csv(f'runs/{run_name}/mnist.csv'), pd.read_csv(f'runs/{run_name}/cifar.csv')
mnist['benchmark'] = 'MNIST'
cifar['benchmark'] = 'CIFAR'

log_data = pd.concat([mnist, cifar])

with open(f'runs/{run_name}/mnist_params.json') as json_file:
Expand All @@ -36,7 +37,12 @@ def load_exp(exp_name):
mnist, cifar = pd.read_csv(f'{exp}/mnist.csv'), pd.read_csv(f'{exp}/cifar.csv')
mnist['benchmark'] = 'MNIST'
cifar['benchmark'] = 'CIFAR'
log_data = pd.concat([mnist, cifar])

svhn, fashion_mnist = pd.read_csv(f'{exp}/svhn.csv'), pd.read_csv(f'{exp}/fashion_mnist.csv')
svhn['benchmark'] = 'SVHN'
fashion_mnist['benchmark'] = 'FashionMNIST'

log_data = pd.concat([mnist, cifar, svhn, fashion_mnist])
log_data[exp_var] = exp_val

exp_dfs.append(log_data)
Expand Down
21 changes: 21 additions & 0 deletions models/fashion_mnist.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,21 @@
import torch
import torch.nn as nn
import torch.nn.functional as F

class Net(nn.Module):
"""Adapted from Pytorch FashionMNIST example"""
def __init__(self):
super(Net, self).__init__()
self.flatten = nn.Flatten()
self.linear_relu_stack = nn.Sequential(
nn.Linear(28*28, 512),
nn.ReLU(),
nn.Linear(512, 512),
nn.ReLU(),
nn.Linear(512, 10),
)

def forward(self, x):
x = self.flatten(x)
logits = self.linear_relu_stack(x)
return logits
42 changes: 42 additions & 0 deletions models/svhn.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
from typing import Tuple, Optional, List

import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader, TensorDataset, Dataset
from torch import Tensor

"""Adapted from Exploring-svhn-using-deep-neural-network, Medium article"""

class Net(nn.Module):
"""Feedfoward neural network with 6 hidden layer"""
def __init__(self):
super().__init__()
# hidden layer
self.linear1 = nn.Linear(32 * 32, 1024)
self.linear2 = nn.Linear(1024, 512)
self.linear3 = nn.Linear(512, 256)
self.linear4 = nn.Linear(256, 128)
self.linear5 = nn.Linear(128, 64)
self.linear6 = nn.Linear(64, 32)
# output layer
self.linear7 = nn.Linear(32, 10)

def forward(self, xb):
# Flatten the image tensors
# xb = xb.view(xb.size(0), -1)
# Get intermediate outputs using hidden layer
out = self.linear1(xb)
out = F.relu(out)
out = self.linear2(out)
out = F.relu(out)
out = self.linear3(out)
out = F.relu(out)
out = self.linear4(out)
out = F.relu(out)
out = self.linear5(out)
out = F.relu(out)
out = self.linear6(out)
out = F.relu(out)
# Get predictions using output layer
out = self.linear7(out)
return out
35 changes: 31 additions & 4 deletions train.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,30 +14,40 @@
from torch import Tensor

import helpers
from models import cifar, mnist
from models import cifar, mnist, svhn, fashion_mnist
from data import utils

from continuum import ClassIncremental
from continuum import InstanceIncremental
from continuum.tasks import split_train_val
from continuum.datasets import MNIST
from continuum.datasets import CIFAR10
from continuum.datasets import SVHN
from continuum.datasets import FashionMNIST

DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')

DATA_ROOT_CIFAR = "data/datasets/cifar-10"
DATA_ROOT_MNIST = "data/datasets/mnist"
DATA_ROOT_SVHN = "data/datasets/svhn"
DATA_ROOT_FashionMNIST = "data/datasets/fashion_mnist"


def load_model(dataset: str):
if dataset == 'cifar':
return cifar.Net().to(DEVICE)
elif dataset == 'mnist':
return mnist.Net().to(DEVICE)
elif dataset == 'svhn':
return svhn.Net().to(DEVICE)
elif dataset == 'fashion_mnist':
return fashion_mnist.Net().to(DEVICE)


def run_continual_exp(exp_name, params, use_devset=False, cl_scenario='Class'):

for benchmark in ['mnist', 'cifar']:
#for benchmark in ['mnist', 'cifar', 'svhn', 'fashion_mnist']:
for benchmark in ['fashion_mnist']:
print(f'Running: {exp_name}/{benchmark}')

if benchmark == 'mnist':
Expand All @@ -48,19 +58,36 @@ def run_continual_exp(exp_name, params, use_devset=False, cl_scenario='Class'):
trainset = MNIST(data_path='data/datasets/mnist', train=True, download=True)
testset = torchvision.datasets.MNIST(DATA_ROOT_MNIST, train=False, download=True, transform=transform)

else:
elif benchmark == 'cifar':
transform = transforms.Compose(
[transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))]
)
trainset = CIFAR10(data_path='data/datasets/cifar10', train=True, download=True)
testset = torchvision.datasets.CIFAR10(DATA_ROOT_CIFAR, train=False, download=True, transform=transform)

elif benchmark == 'svhn':
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,))
])
trainset = SVHN(data_path='data/datasets/svhn', train=True, download=True)
testset = torchvision.datasets.SVHN(DATA_ROOT_SVHN, split='train', download=True, transform=transform)

else: #benchmark == 'fashion_mnist':
transform = transforms.Compose([
transforms.Resize(32),
transforms.ToTensor()
])
trainset = FashionMNIST(data_path='data/datasets/fashion_mnist', train=True, download=True)
testset = torchvision.datasets.FashionMNIST(DATA_ROOT_FashionMNIST, train=False, download=True, transform=transform)

if cl_scenario == 'Class':
scenario = ClassIncremental(trainset, transformations=[transform], increment=1)
print(f"Number of classes: {scenario.nb_classes}.")
print(f"Number of tasks: {scenario.nb_tasks}.")
else:
scenario = InstanceIncremental(dataset=trainset, transformations=[transform], nb_tasks=10)
scenario = InstanceIncremental(trainset, transformations=[transform], nb_tasks=10)
# scenario = InstanceIncremental(dataset=trainset, transformations=[transform], nb_tasks=10)
print(f"Number of tasks: {scenario.nb_tasks}")

model = load_model(dataset=benchmark)
Expand Down