-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathExpt.py
More file actions
127 lines (117 loc) · 5 KB
/
Copy pathExpt.py
File metadata and controls
127 lines (117 loc) · 5 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
114
115
116
117
118
119
120
121
122
123
124
125
126
127
"""Train Microbiome classifiers with Random Forests."""
import argparse
from utils.training import intra_dataset, cross_dataset, LODO
# From https://sumit-ghosh.com/articles/parsing-dictionary-key-value-pairs-kwargs-argparse-python/
class ParseKwargs(argparse.Action):
def __call__(self, parser, namespace, values, option_string=None):
setattr(namespace, self.dest, dict())
for value in values:
key, value = value.split('=')
if value.isdigit():
value = int(value)
elif value == 'None':
value = None
else:
try:
value = float(value)
except:
pass
getattr(namespace, self.dest)[key] = value
# Parse command line arguments
parser = argparse.ArgumentParser()
parser.add_argument('--analysis_type',
nargs='+',
default='all',
help='intra, cross, lodo, or all')
parser.add_argument('-dn', '--data_name',
help='name of dataset, should be saved as data/data_name.pkl')
parser.add_argument('--file_suffix',
default='',
help='suffix for results file')
parser.add_argument('-nr', '--num_repetitions',
type=int,
default=50,
help='number of times to repeat experiment')
parser.add_argument('-nf', '--num_folds',
type=int,
default=5,
help='number of folds for k-fold cross-validation')
parser.add_argument('-tt', '--transformation_type',
choices=['prop', 'clr', 'alr', 'none'],
default='prop',
help='data transformation')
parser.add_argument('-to', '--transformation_opts',
nargs='+',
action=ParseKwargs,
help='data transformation options (as key=value)')
parser.add_argument('-ne', '--n_estimators',
type=int,
default=500,
help='n_estimators in scikit-learn RandomForestClassifier')
parser.add_argument('-mf', '--max_features',
default='sqrt',
help='max_features in scikit-learn RandomForestClassifier')
parser.add_argument('-md', '--max_depth',
default=None,
help='max_depth in scikit-learn RandomForestClassifier')
parser.add_argument('-ms', '--max_samples',
default=None,
help='max_samples in scikit-learn RandomForestClassifier')
parser.add_argument('-mss', '--min_samples_split',
default=2,
help='min_samples_split in scikit-learn RandomForestClassifier')
parser.add_argument('-im', '--importance_measure',
choices=['SHAP', 'permutation'],
default='SHAP',
help='measure of feature importance')
parser.add_argument('-lb', '--balanced',
action='store_true',
help='use balanced samples for LODO')
args = parser.parse_args()
# Convert max_features to integer, float, or string
if args.max_features.isdigit():
max_features = int(args.max_features)
else:
try:
max_features = float(args.max_features)
except ValueError:
max_features = args.max_features
# Make dictionary of transformation options
if args.transformation_opts is None:
transformation_opts = {}
else:
transformation_opts = args.transformation_opts
# Dictionary of model hyperparameters
model_opts = {"n_estimators": args.n_estimators,
"max_features": max_features,
"max_depth": args.max_depth,
"max_samples": args.max_samples,
"min_samples_split": args.min_samples_split}
if 'all' in args.analysis_type:
analysis_type = ['intra', 'cross', 'lodo']
else:
analysis_type = args.analysis_type
if 'intra' in analysis_type:
intra_dataset(data_name=args.data_name,
file_suffix=args.file_suffix,
num_repetitions=args.num_repetitions,
num_folds=args.num_folds,
transformation_type=args.transformation_type,
transformation_opts=transformation_opts,
model_opts=model_opts)
if 'cross' in analysis_type:
cross_dataset(data_name=args.data_name,
file_suffix=args.file_suffix,
num_repetitions=args.num_repetitions,
transformation_type=args.transformation_type,
transformation_opts=transformation_opts,
model_opts=model_opts,
importance_measure=args.importance_measure)
if 'lodo' in analysis_type:
LODO(data_name=args.data_name,
file_suffix=args.file_suffix,
num_repetitions=args.num_repetitions,
transformation_type=args.transformation_type,
transformation_opts=transformation_opts,
model_opts=model_opts,
balanced=args.balanced)