-
Notifications
You must be signed in to change notification settings - Fork 3
Expand file tree
/
Copy pathmain.py
More file actions
73 lines (58 loc) · 2.9 KB
/
Copy pathmain.py
File metadata and controls
73 lines (58 loc) · 2.9 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
import logging
import os
from pathlib import Path
import pickle
import click
import torch
import adversarial_experiment
import configurations
import plots
import taylor_experiment
import utils
@click.command()
@click.option('--taylor-exp', type=str, default=None, help='Should run the Taylor expansion experiment.')
@click.option('--adversarial-exp', type=str, default=None, help='Should run the adversarial accuracy experiment.')
def main(taylor_exp: str, adversarial_exp: str):
""" Main function that runs either the taylor or the adversarial experiment.
:param taylor_exp: str, type of taylor experiment to run. Must be 'test' or 'final'
:param adversarial_exp: str, type of adversarial experiment to run. Must be 'test' or 'spirals_penalization'
:return:None
"""
ROOT_DIR = str(Path(__file__).resolve().parents[0])
if taylor_exp is not None:
print('Running the Taylor expansion experiment')
if taylor_exp not in configurations.configs_taylor_exp:
raise ValueError('--taylor-exp should be among {}'.format(list(configurations.configs_taylor_exp.keys())))
config = configurations.configs_taylor_exp[taylor_exp]
save_dir = os.path.join(ROOT_DIR, 'results', 'taylor', taylor_exp)
utils.clean_dir(save_dir)
pickle.dump(config, open(os.path.join(save_dir, 'config.pkl'), 'wb'))
print('Computing Taylor convergence')
taylor_experiment.compute_taylor_convergence(save_dir, config)
print('Plotting results')
logging.disable(logging.WARNING) # Disable matplotlib warnings
plots.plot_taylor_convergence(save_dir)
logging.disable(logging.NOTSET) # Re-enable warnings
if adversarial_exp is not None:
print('Running the adversarial accuracy experiment')
if adversarial_exp not in configurations.configs_adversarial_exp:
raise ValueError('--adversarial-exp should be among {}'.format(list(configurations.configs_adversarial_exp.keys())))
config = configurations.configs_adversarial_exp[adversarial_exp]
save_dir = os.path.join(ROOT_DIR, 'results', 'adversarial', adversarial_exp)
utils.clean_dir(save_dir)
config['save_dir'] = [save_dir]
print('Training model')
utils.gridsearch(adversarial_experiment.ex, config, save_dir)
print('Computing adversarial accuracy')
adversarial_experiment.compute_adversarial_accuracy(save_dir)
print('Computing training norms')
adversarial_experiment.compute_norms(save_dir, run_nums=['1', '2'])
print('Plotting results')
logging.disable(logging.WARNING) # Disable matplotlib warnings
plots.plot_spirals_adversarial(save_dir)
plots.plot_spirals_training_norms(save_dir)
logging.disable(logging.NOTSET) # Re-enable warnings
if __name__ == '__main__':
# Check if GPU is used
print('GPU available: ', torch.cuda.is_available())
main()