Skip to content

Commit e65e588

Browse files
author
Colin Grambow
committed
Build entire model before predicting
1 parent 73edb8f commit e65e588

2 files changed

Lines changed: 64 additions & 6 deletions

File tree

reacdiff/parsing.py

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,16 +15,53 @@ def parse_predict_args():
1515
parser = argparse.ArgumentParser(
1616
formatter_class=argparse.ArgumentDefaultsHelpFormatter
1717
)
18+
19+
# General arguments
1820
parser.add_argument('--data_path', type=str, required=True,
1921
help='Path to data containing states for prediction task')
2022
parser.add_argument('--model', type=str, required=True,
2123
help='Path to trained model')
24+
parser.add_argument('--output_dim', type=int, required=True,
25+
help='Dimensionality of predictions')
2226
parser.add_argument('--data_path2', type=str,
2327
help='Path to additional observable states for prediction')
2428
parser.add_argument('--save_path', type=str, default=os.path.join(os.getcwd(), 'preds.csv'),
2529
help='Path to save predictions to')
30+
parser.add_argument('--cpu', action='store_true',
31+
help='Run on CPU instead of GPU')
2632
parser.add_argument('--batch_size', type=int, default=32,
2733
help='Batch size')
34+
35+
# Encoder arguments
36+
parser.add_argument('--feat_maps', type=int, default=16,
37+
help='Number of feature maps in first convolutional layer')
38+
parser.add_argument('--first_conv_size', type=int, default=7,
39+
help='Filter size in first convolutional layer')
40+
parser.add_argument('--first_conv_stride', type=int, default=2,
41+
help='Strides in first convolutional layer')
42+
parser.add_argument('--first_pool_size', type=int, default=3,
43+
help='Window size in first pool')
44+
parser.add_argument('--first_pool_stride', type=int, default=2,
45+
help='Strides in first pool')
46+
parser.add_argument('--growth_rate', type=int, default=12,
47+
help='Growth rate in dense blocks')
48+
parser.add_argument('--blocks', type=int, nargs='+', default=[3, 4, 5],
49+
help='Numbers of layers in each dense block')
50+
parser.add_argument('--dropout', type=float, default=0.0,
51+
help='Dropout rate after convolutions')
52+
parser.add_argument('--reduction', type=float, default=0.5,
53+
help='Compression rate in transition blocks')
54+
parser.add_argument('--no_bottleneck', action='store_true',
55+
help='Do not use bottleneck convolution in dense blocks')
56+
parser.add_argument('--flatten_last', action='store_true',
57+
help='Flatten instead of global pool before output')
58+
59+
# RNN arguments
60+
parser.add_argument('--rnn_layers', type=int, default=1,
61+
help='Number of RNN layers')
62+
parser.add_argument('--rnn_units', type=int, default=100,
63+
help='Number of units in RNN')
64+
2865
return parser.parse_args()
2966

3067

reacdiff/train/predict.py

Lines changed: 27 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,7 @@
11
import os
22

3-
import keras
4-
53
import reacdiff.data.data as datamod
6-
import reacdiff.utils as utils
4+
import reacdiff.models.crnn as models
75

86

97
def predict(args):
@@ -16,9 +14,32 @@ def predict(args):
1614

1715
os.makedirs(os.path.dirname(args.save_path), exist_ok=True)
1816

19-
# Load model
20-
model = keras.models.load_model(args.model, custom_objects={'rmse': utils.rmse, 'mae': utils.mae})
17+
# Load model (currently have to build entire model with provided arguments)
18+
# Build model
19+
crnn = models.CRNN(
20+
time_steps=data.data.shape[1],
21+
input_shape=data.data.shape[2:],
22+
output_dim=args.output_dim,
23+
observables=data.get_num_observables(),
24+
rnn_layers=args.rnn_layers,
25+
rnn_units=args.rnn_units,
26+
use_gpu=not args.cpu
27+
)
28+
crnn.build(
29+
feat_maps=args.feat_maps,
30+
first_conv_size=args.first_conv_size,
31+
first_conv_stride=args.first_conv_stride,
32+
first_pool_size=args.first_pool_size,
33+
first_pool_stride=args.first_pool_stride,
34+
growth_rate=args.growth_rate,
35+
blocks=args.blocks,
36+
dropout=args.dropout,
37+
reduction=args.reduction,
38+
bottleneck=not args.no_bottleneck,
39+
flatten_last=args.flatten_last
40+
)
41+
crnn.model.load_weights(args.model)
2142

2243
# Predict
23-
preds = model.predict(data.get_data(), batch_size=args.batch_size, verbose=1)
44+
preds = crnn.model.predict(data.get_data(), batch_size=args.batch_size, verbose=1)
2445
datamod.save_csv(preds, args.save_path)

0 commit comments

Comments
 (0)