-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathmain_visualization.py
More file actions
84 lines (62 loc) · 3.67 KB
/
Copy pathmain_visualization.py
File metadata and controls
84 lines (62 loc) · 3.67 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
import pdb
import argparse
import sys
import os
import viz.viz_functions as vis_tools
sys.path.append(os.getcwd())
parser = argparse.ArgumentParser()
parser.add_argument("--viz_method", type=str, default="none", choices = ["pca", "tsne", "none"])
parser.add_argument("--features_to_show", type=str, default="meds_diags_procs", choices = ["meds_diags_procs","all"])
parser.add_argument("--sampled", type=int, default=1, choices = [0,1])
parser.add_argument("--sample_size", type=int, default=1000)
parser.add_argument("--perplex", type=int, default=35)
parser.add_argument("--num_it", type=int, default=1000)
parser.add_argument("--lr_rate", type=int, default=200)
parser.add_argument("--sample_size_for_shap", type=float, default=0.05)
parser.add_argument("--trained_model_path", type=str, default="saved_classical_ml_models/rf_model.pkl")
parser.add_argument("--compute_table_1", type=int, default=0, choices = [0, 1])
parser.add_argument("--plot_violins_flag", type=int, default=0, choices = [0, 1])
parser.add_argument("--train_stationary_filename", type=str, default="stationary_data/stationary_data_imbratio1_normalized_train.csv")
parser.add_argument("--test_stationary_filename", type=str, default="stationary_data/stationary_data_imbratio1_normalized_test.csv")
parser.add_argument("--feature_ranking_path", type=str, default="saved_classical_ml_models/feature_impoerance_rf.csv")
parser.add_argument("--mci_metadata", type=str, default="intermediate_files/mci_metadata.csv")
parser.add_argument("--nonmci_metadata", type=str, default="intermediate_files/nonmci_metadata.csv")
if parser.parse_args().viz_method == "tsne":
args = parser.parse_args()
vis_tools.tSNE_visualization(args.train_stationary_filename
, args.test_stationary_filename
, args.sampled
, args.sample_size
, args.features_to_show
, args.perplex
, args.num_it
, args.lr_rate)
elif parser.parse_args().viz_method == "pca":
args = parser.parse_args()
vis_tools.pca_visualization(args.train_stationary_filename
, args.test_stationary_filename
, args.sampled
, args.sample_size
, args.features_to_show
)
elif parser.parse_args().viz_method == "none":
print("Warning: no visualization method has been selected.")
if parser.parse_args().compute_table_1 == 1:
args = parser.parse_args()
vis_tools.compute_table_stats(args.train_stationary_filename
, args.test_stationary_filename
, args.features_to_show
, args.mci_metadata
, args.nonmci_metadata
, args.feature_ranking_path
)
elif parser.parse_args().compute_table_1 == 0:
print("Warning: no compute_table_1 method has been selected.")
if parser.parse_args().plot_violins_flag == 1:
args = parser.parse_args()
vis_tools.plot_violins(args.train_stationary_filename
, args.test_stationary_filename
, args.feature_ranking_path
, args.sampled
, args.sample_size
)