-
Notifications
You must be signed in to change notification settings - Fork 3
Expand file tree
/
Copy pathrun.py
More file actions
46 lines (32 loc) · 1.54 KB
/
Copy pathrun.py
File metadata and controls
46 lines (32 loc) · 1.54 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
import os
import sys
import argparse
import datetime
from models.model_loader import load_model
from data.dataset_loader import load_data
from utils.utils import initialize_environment, run, parse_config
from evaluation.evaluator import Evaluator
os.environ['MKL_NUM_THREADS'] = "1"
def main(config):
initialize_environment(config)
model, config = load_model(config)
sensitive_train_loader, sensitive_val_loader, sensitive_test_loader, public_train_loader, config = load_data(config)
model.pretrain(public_train_loader, config.pretrain)
model.train(sensitive_train_loader, config.train)
syn_data, syn_labels = model.generate(config.gen)
evaluator = Evaluator(config)
evaluator.eval(syn_data, syn_labels, sensitive_train_loader, sensitive_val_loader, sensitive_test_loader)
# evaluator.eval_fidelity(syn_data, syn_labels, sensitive_train_loader, sensitive_val_loader, sensitive_test_loader)
if __name__ == '__main__':
sys.path.append(os.getcwd())
parser = argparse.ArgumentParser()
parser.add_argument('--config_dir', default="configs")
parser.add_argument('--method', '-m', default="DP-LDM")
parser.add_argument('--epsilon', '-e', default="10.0")
parser.add_argument('--data_name', '-dn', default="cifar10_32")
parser.add_argument('--exp_description', '-ed', default="")
parser.add_argument('--resume_exp', '-re', default=None)
parser.add_argument('--config_suffix', '-cs', default="")
opt, unknown = parser.parse_known_args()
config = parse_config(opt, unknown)
run(main, config)