-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathargparser.py
More file actions
113 lines (107 loc) · 4.59 KB
/
Copy pathargparser.py
File metadata and controls
113 lines (107 loc) · 4.59 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
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
import argparse
import os
dir_prefix = os.getcwd()
coat_params = {
"train_path": dir_prefix + "/data_process/coat/train.csv",
"random_path": dir_prefix + "/data_process/coat/random.csv",
"user_feature_path": dir_prefix + "/data_process/coat/user_feat_onehot.csv",
"user_feature_label": dir_prefix + "/data_process/coat/user_feat_label.csv",
"dcf_A_hat_path": dir_prefix + "/data_process/coat_wg_Afit/",
"vae_path": dir_prefix + "/data_process/coat_vae/",
"ivae_path": dir_prefix + "/data_process/coat_ivae/",
"test_ratio": 0.7,
"train_ratio": 0.8,
"user_feature_dim": [2, 6, 3, 3, 2, 16, 13, 2],
"threshold": 4.0,
"min_val": 1.0,
"max_val": 5.0,
"batch_size": 512,
"beta_max": 1.,
"name": "coat",
}
yahoo_params = {
"train_path": dir_prefix + "/data_process/Yahoo_R3/train.csv",
"random_path": dir_prefix + "/data_process/Yahoo_R3/random.csv",
"user_feature_path": dir_prefix + "/data_process/Yahoo_R3/user_feat_onehot.csv",
"user_feature_label": dir_prefix + "/data_process/Yahoo_R3/user_feat_label.csv",
"dcf_A_hat_path": dir_prefix + "/data_process/R3_wg_Afit/",
"ivae_path": dir_prefix + "/data_process/yahoo_ivae/",
"vae_path": dir_prefix + "/data_process/yahoo_vae/",
"test_ratio": 0.7,
"train_ratio": 0.8,
"user_feature_dim": [5, 5, 5, 5, 5, 5, 5],
"threshold": 4.0,
"min_val": 1.0,
"max_val": 5.0,
"batch_size": 512,
"beta_max": 1.,
"name": "yahoo"
}
kuai_rand_params = {
"train_path": dir_prefix + "/data_process/kuai_rand/train.csv",
"random_path": dir_prefix + "/data_process/kuai_rand/random.csv",
"user_feature_path": dir_prefix + "/data_process/kuai_rand/user_feat_onehot.csv",
"user_feature_label": dir_prefix + "/data_process/kuai_rand/user_feat_label.csv",
"dcf_A_hat_path": dir_prefix + "/data_process/kuai_rand_wg_Afit/",
"ivae_path": dir_prefix + "/data_process/kuai_rand_ivae/",
"vae_path": dir_prefix + "/data_process/kuai_rand_vae/",
"test_ratio": 0.7,
"train_ratio": 0.8,
"user_feature_dim": [9, 2, 2, 8, 9, 7, 8, 2, 7, 50, 1471, 33, 3, 118, 454, 7, 5, 4],
"threshold": 0.9,
"min_val": 0.0,
"max_val": 5.0,
"batch_size": 2048,
"beta_max": 1.,
"name": "kuai_rand"
}
simulation_params = {
"train_path": dir_prefix + "/data_process/simulation/train.csv",
"random_path": dir_prefix + "/data_process/simulation/random.csv",
"user_feature_path": dir_prefix + "/data_process/simulation/user_feat_onehot.csv",
"user_feature_label": dir_prefix + "/data_process/simulation/user_feat_label.csv",
"dcf_A_hat_path": dir_prefix + "/data_process/sim_wg_Afit/",
"ivae_path": dir_prefix + "/data_process/sim_ivae/",
"vae_path": dir_prefix + "/data_process/sim_vae/",
"test_ratio": 0.7,
"train_ratio": 0.8,
"user_feature_dim": [5, ],
"threshold": 4.0,
"min_val": 0.0,
"max_val": 5.0,
"batch_size": 1024,
"beta_max": 1.,
"name": "sim"
}
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument("--dir_prefix", type=str, default=dir_prefix)
parser.add_argument("--tune", action="store_true")
parser.add_argument("--metric", type=str, default="ndcg")
parser.add_argument("--dataset", type=str, default="coat")
parser.add_argument("--patience", type=int, default=5)
parser.add_argument("--topk", type=int, default=5)
parser.add_argument("--seed", type=int, default=1234)
parser.add_argument("--test_seed", action="store_true")
parser.add_argument("--sim_suffix", type=str, default="")
parser.add_argument("--key_name", type=str)
args = parser.parse_args()
if args.dataset == "yahoo":
data_params = yahoo_params
elif args.dataset == "coat":
data_params = coat_params
elif args.dataset == "kuai_rand":
data_params = kuai_rand_params
elif args.dataset == "sim":
data_params = simulation_params
data_params["train_path"] = dir_prefix + "/data_process/simulation/train{}.csv".format(args.sim_suffix)
data_params["random_path"] = dir_prefix + "/data_process/simulation/random{}.csv".format(args.sim_suffix)
sr = args.sim_suffix.split("_")[2]
tr = args.sim_suffix.split("_")[-1]
data_params["ivae_path"] = dir_prefix + "/data_process/sim_ivae/sr_{}_tr_{}/".format(sr, tr)
data_params["vae_path"] = dir_prefix + "/data_process/sim_vae/sr_{}_tr_{}/".format(sr, tr)
data_params["dcf_A_hat_path"] = dir_prefix + "/data_process/sim_wg_Afit/sr_{}_tr_{}/".format(sr, tr)
else:
raise Exception("invalid dataset")
setattr(args, "data_params", data_params)
return args