-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathNamespace.py
More file actions
79 lines (72 loc) · 2.97 KB
/
Copy pathNamespace.py
File metadata and controls
79 lines (72 loc) · 2.97 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
74
75
import torch
class Namespace:
def __init__(self, settings, experiment):
self.data = settings['data']
self.device = torch.device('cuda:' + str(settings['gpu_id']) if torch.cuda.is_available() else 'cpu')
self.backbone = settings['baseline']
self.mo_method = settings['wrapper']
self.mode = experiment['mode']
self.every = settings['validation_rate']
self.metric = settings['validation_metric']
self.batch_size = settings['batch_size']
self.n_epochs = settings['epochs']
try:
self.seed = experiment['seed']
except KeyError:
self.seed = 42
if self.backbone == 'BPRMF':
self.dim = experiment['dim']
self.lr = experiment['lr']
self.weight_decay = experiment['l_2']
elif self.backbone == 'DirectAU':
self.dim = experiment['dim']
self.lr = experiment['lr']
self.weight_decay = experiment['l_2']
self.gamma = experiment['gamma']
self.patience = experiment['patience']
elif self.backbone == 'LightGCN':
self.dim = experiment['dim']
self.lr = experiment['lr']
self.weight_decay = experiment['l_2']
self.layers = experiment['layers']
self.normalize = experiment['normalize']
elif self.backbone == 'NGCF':
self.dim = experiment['dim']
self.lr = experiment['lr']
self.weight_decay = experiment['l_2']
self.layers = experiment['layers']
self.message_dropout = experiment['message_dropout']
self.node_dropout = experiment['node_dropout']
self.normalize = experiment['normalize']
if self.mo_method == 'AMORE_MGDA':
self.atk = experiment['atk']
self.type = experiment['g_n']
self.ranker = experiment['ranker']
try:
self.atk_con = experiment['atk']['atk_cons']
except KeyError:
pass
try:
self.atk_pro = experiment['atk']['atk_prov']
except KeyError:
pass
if self.ranker == 'base':
self.ablation = experiment['ablation']
elif self.mo_method == 'multifr':
self.gamma = experiment['gamma']
self.temp = experiment['temp']
self.type = experiment['g_n']
self.ranker = experiment['ranker']
elif self.mo_method == 'None':
self.scale1 = experiment['scale']
elif self.mo_method in ['AMORE_SCALE', 'AMORE_ABL', 'AMORE_ABL_WOS', 'AMORE_ABL_WOZ', 'AMORE_EPO']:
try:
self.atk_con = experiment['atk']['atk_cons']
except KeyError:
pass
try:
self.atk_pro = experiment['atk']['atk_prov']
except KeyError:
pass
self.ranker = experiment['ranker']
self.scale1 = experiment['scale']