-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathtrain.py
More file actions
50 lines (40 loc) · 1.64 KB
/
Copy pathtrain.py
File metadata and controls
50 lines (40 loc) · 1.64 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
import sklearn # do this first, otherwise get a libgomp error?!
import argparse, os, sys, random, logging
import numpy as np
from sklearn.linear_model import LogisticRegression
import pickle
# Set random seeds for repeatable results
RANDOM_SEED = 3
random.seed(RANDOM_SEED)
np.random.seed(RANDOM_SEED)
# Load files
parser = argparse.ArgumentParser(description='Train custom ML model')
parser.add_argument('--data-directory', type=str, required=True)
parser.add_argument('--epochs', type=int, required=True)
parser.add_argument('--out-directory', type=str, required=True)
args, _ = parser.parse_known_args()
out_directory = args.out_directory
if not os.path.exists(out_directory):
os.mkdir(out_directory)
# grab train/test set
X_train = np.load(os.path.join(args.data_directory, 'X_split_train.npy'))
Y_train = np.load(os.path.join(args.data_directory, 'Y_split_train.npy'))
X_test = np.load(os.path.join(args.data_directory, 'X_split_test.npy'))
Y_test = np.load(os.path.join(args.data_directory, 'Y_split_test.npy'))
# sparse representation of the labels (1-based)
Y_train = np.argmax(Y_train, axis=1) + 1
Y_test = np.argmax(Y_test, axis=1) + 1
print('Training model on', str(X_train.shape[0]), 'inputs...')
# train your model
clf = LogisticRegression(random_state=RANDOM_SEED, max_iter=args.epochs)
clf.fit(X_train, Y_train)
print('Training model OK')
print('')
print('Mean accuracy (training set):', clf.score(X_train, Y_train))
print('Mean accuracy (validation set):', clf.score(X_test, Y_test))
print('')
print('Saving model...')
with open(os.path.join(args.out_directory, 'model.pkl'),'wb') as f:
pickle.dump(clf, f)
print('Saving model OK')
print('')