()
| 32 | |
| 33 | |
| 34 | def test_keras_traveller(): |
| 35 | kmt = utils.KerasModelTraveller() |
| 36 | model = sam_32_32() |
| 37 | layer_names = kmt.get_layer_names(keras_model=model) |
| 38 | expected_layer_names = [ |
| 39 | "input_1", |
| 40 | "conv2d", |
| 41 | "re_lu", |
| 42 | "conv2d_1", |
| 43 | "re_lu_1", |
| 44 | "conv2d_2", |
| 45 | "re_lu_2", |
| 46 | "conv2d_3", |
| 47 | "add", |
| 48 | "re_lu_3", |
| 49 | "conv2d_4", |
| 50 | "re_lu_4", |
| 51 | "conv2d_5", |
| 52 | "add_1", |
| 53 | "re_lu_5", |
| 54 | "conv2d_6", |
| 55 | "re_lu_6", |
| 56 | "conv2d_7", |
| 57 | "conv2d_8", |
| 58 | "add_2", |
| 59 | "re_lu_7", |
| 60 | "max_pooling2d", |
| 61 | "flatten", |
| 62 | "dense", |
| 63 | "re_lu_8", |
| 64 | "dense_1", |
| 65 | ] |
| 66 | assert layer_names == expected_layer_names, "Keras model traveller failed." |
| 67 | tf.keras.backend.clear_session() |
| 68 | |
| 69 | |
| 70 | def test_convert_to_onnx(): |
nothing calls this directly
no test coverage detected