-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathNN-RHC.py
More file actions
50 lines (44 loc) · 1.97 KB
/
Copy pathNN-RHC.py
File metadata and controls
50 lines (44 loc) · 1.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
"""
RHC NN training on HTRU2 data
"""
# Adapted from https://github.com/JonathanTay/CS-7641-assignment-2/blob/master/NN1.py
import sys
sys.path.append("./ABAGAIL/ABAGAIL.jar")
from func.nn.backprop import BackPropagationNetworkFactory
from shared import SumOfSquaresError, DataSet, Instance
from opt.example import NeuralNetworkOptimizationProblem
from func.nn.backprop import RPROPUpdateRule
import opt.RandomizedHillClimbing as RandomizedHillClimbing
from func.nn.activation import RELU
from base import *
# Network parameters found "optimal" in Assignment 1
INPUT_LAYER = 23
HIDDEN_LAYER1 = 32
HIDDEN_LAYER2 = 32
OUTPUT_LAYER = 1
TRAINING_ITERATIONS = 20001
OUTFILE = OUTPUT_DIRECTORY + '/NN_OUTPUT/NN_{}_LOG.csv'
def main():
"""Run this experiment"""
training_ints = initialize_instances(TRAIN_DATA_FILE)
testing_ints = initialize_instances(TEST_DATA_FILE)
validation_ints = initialize_instances(VALIDATE_DATA_FILE)
factory = BackPropagationNetworkFactory()
measure = SumOfSquaresError()
data_set = DataSet(training_ints)
relu = RELU()
# 50 and 0.000001 are the defaults from RPROPUpdateRule.java
rule = RPROPUpdateRule(0.064, 50, 0.000001)
oa_names = ["RHC"]
classification_network = factory.createClassificationNetwork(
[INPUT_LAYER, HIDDEN_LAYER1, HIDDEN_LAYER2, OUTPUT_LAYER], relu)
nnop = NeuralNetworkOptimizationProblem(data_set, classification_network, measure)
oa = RandomizedHillClimbing(nnop)
train(oa, classification_network, 'RHC', training_ints, validation_ints, testing_ints, measure,
TRAINING_ITERATIONS, OUTFILE.format('RHC'))
if __name__ == "__main__":
with open(OUTFILE.format('RHC'), 'a+') as f:
f.write('{},{},{},{},{},{},{},{},{},{},{}\n'.format('iteration', 'MSE_trg', 'MSE_val', 'MSE_tst', 'acc_trg',
'acc_val', 'acc_tst', 'f1_trg', 'f1_val', 'f1_tst',
'elapsed'))
main()