forked from deepsound-project/genre-recognition
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmodel_to_tfjs.py
More file actions
31 lines (27 loc) · 1.21 KB
/
Copy pathmodel_to_tfjs.py
File metadata and controls
31 lines (27 loc) · 1.21 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
from tensorflow.keras.models import Model, load_model
import tensorflowjs as tfjs
from optparse import OptionParser
import os
def extract_realtime_model(full_model):
input = full_model.get_layer('input').input
output = full_model.get_layer('output_realtime').output
model = Model(inputs=input, outputs=output)
return model
def main(model_path, output_path):
model = load_model(model_path)
realtime_model = extract_realtime_model(model)
realtime_model.compile(optimizer=model.optimizer, loss=model.loss)
tfjs.converters.save_keras_model(realtime_model, output_path)
if __name__ == '__main__':
parser = OptionParser()
parser.add_option('-m', '--model_path', dest='model_path',
default=os.path.join(os.path.dirname(__file__),
'models/model.h5'),
help='path to the input model YAML file', metavar='MODEL_PATH')
parser.add_option('-o', '--output_path', dest='output_path',
default=os.path.join(os.path.dirname(__file__),
'static/model'),
help='path to the output TFJS model directory',
metavar='OUTPUT_PATH')
options, args = parser.parse_args()
main(options.model_path, options.output_path)