Add realistic model input for mnist and option for float model input range Added realistic model input image from mnist. As mnist input range is different from others, added a general option for float model input range. Change-Id: Ib67e0fe0d1a4a52f475ad24978db305cd1da0161
diff --git a/build_tools/gen_mlmodel_input.py b/build_tools/gen_mlmodel_input.py index 5a910a1..117aacb 100755 --- a/build_tools/gen_mlmodel_input.py +++ b/build_tools/gen_mlmodel_input.py
@@ -19,6 +19,8 @@ help='Model input shape (example: "1, 224, 224, 3")', required=True) parser.add_argument('--q', dest='is_quant', action='store_true', help='Indicate it is quant model (default: False)') +parser.add_argument('--r', dest='float_input_range', default="-1.0, 1.0", + help='Float model input range (default: "-1.0, 1.0")') args = parser.parse_args() @@ -40,12 +42,15 @@ (input_shape[1], input_shape[2])) input = np.array(resized_img).reshape(np.prod(input_shape)) if not is_quant: - input = 2.0 / 255.0 * input - 1 + low = np.min(float_input_range) + high = np.max(float_input_range) + input = (high - low) * input / 255.0 + low write_binary_file(output_file, input, is_quant) if __name__ == '__main__': # convert input shape to a list input_shape = [int(x) for x in args.input_shape.split(',')] + float_input_range = [float(x) for x in args.float_input_range.split(',')] gen_mlmodel_input(args.input_name, args.output_file, input_shape, args.is_quant)
diff --git a/cmake/iree_model_input.cmake b/cmake/iree_model_input.cmake index f9e09ca..7be9e35 100644 --- a/cmake/iree_model_input.cmake +++ b/cmake/iree_model_input.cmake
@@ -27,7 +27,7 @@ cmake_parse_arguments( _RULE "QUANT" - "NAME;SHAPE;SRC" + "NAME;SHAPE;SRC;RANGE" "" ${ARGN} ) @@ -56,6 +56,9 @@ list(APPEND _ARGS "--i=${_INPUT_FILENAME}") list(APPEND _ARGS "--o=${_OUTPUT_BINARY}") list(APPEND _ARGS "--s=${_RULE_SHAPE}") + if(_RULE_RANGE) + list(APPEND _ARGS "--r=${_RULE_RANGE}") + endif() if(_RULE_QUANT) list(APPEND _ARGS "--q") endif()
diff --git a/samples/float_model_embedding/CMakeLists.txt b/samples/float_model_embedding/CMakeLists.txt index fbd36b3..029609e 100644 --- a/samples/float_model_embedding/CMakeLists.txt +++ b/samples/float_model_embedding/CMakeLists.txt
@@ -158,6 +158,18 @@ example_images/YellowLabradorLooking_new.jpg" ) +iree_model_input( + NAME + mnist_input + SHAPE + "1, 28, 28, 1" + SRC + "https://raw.githubusercontent.com/google/iree/main/ \ + iree/samples/vision/mnist_test.png" + RANGE + "0, 1" +) + if(${BUILD_INTERNAL_MODELS}) iree_model_input( @@ -206,6 +218,7 @@ "mnist.c" DEPS ::mnist_bytecode_module_dylib_c + ::mnist_input_c samples::util::util LINKOPTS "LINKER:--defsym=__stack_size__=100k"
diff --git a/samples/float_model_embedding/mnist.c b/samples/float_model_embedding/mnist.c index 3e7d121..5c753a3 100644 --- a/samples/float_model_embedding/mnist.c +++ b/samples/float_model_embedding/mnist.c
@@ -10,6 +10,7 @@ // Compiled module embedded here to avoid file IO: #include "samples/float_model_embedding/mnist_bytecode_module_dylib_c.h" +#include "samples/float_model_embedding/mnist_input_c.h" const MlModel kModel = { .num_input = 1, @@ -34,14 +35,8 @@ iree_status_t load_input_data(const MlModel *model, void **buffer) { iree_status_t result = alloc_input_buffer(model, buffer); - // Populate initial value - srand(74886738); - if (iree_status_is_ok(result)) { - for (int i = 0; i < model->input_length[0]; ++i) { - int x = rand(); - ((float *)*buffer)[i] = (float)x / (float)RAND_MAX; - } - } + const struct iree_file_toc_t *input_file_toc = mnist_input_c_create(); + memcpy(*buffer, input_file_toc->data, input_file_toc->size); return result; } @@ -56,5 +51,16 @@ } LOG_INFO("Output #%d data length: %d", index_output, mapped_memory->contents.data_length / model->output_size_bytes); + // find the label index with best prediction + float best_out = 0.0; + int best_idx = -1; + for (int i = 0; i < model->output_length[index_output]; ++i) { + float out = ((float *)mapped_memory->contents.data)[i]; + if (out > best_out) { + best_out = out; + best_idx = i; + } + } + LOG_INFO("Digit recognition result is: digit: %d", best_idx); return result; }