| 26 | namespace { |
| 27 | |
| 28 | TEST(ModelCmdlineFlagsTest, ParseArgsStringMapList) { |
| 29 | int args_count = 3; |
| 30 | const char* args[] = { |
| 31 | "toco", "--input_arrays=input_1", |
| 32 | "--rnn_states={state_array:rnn/BasicLSTMCellZeroState/zeros," |
| 33 | "back_edge_source_array:rnn/basic_lstm_cell/Add_1,size:4}," |
| 34 | "{state_array:rnn/BasicLSTMCellZeroState/zeros_1," |
| 35 | "back_edge_source_array:rnn/basic_lstm_cell/Mul_2,size:4}", |
| 36 | nullptr}; |
| 37 | |
| 38 | string expected_input_arrays = "input_1"; |
| 39 | std::vector<std::unordered_map<string, string>> expected_rnn_states; |
| 40 | expected_rnn_states.push_back( |
| 41 | {{"state_array", "rnn/BasicLSTMCellZeroState/zeros"}, |
| 42 | {"back_edge_source_array", "rnn/basic_lstm_cell/Add_1"}, |
| 43 | {"size", "4"}}); |
| 44 | expected_rnn_states.push_back( |
| 45 | {{"state_array", "rnn/BasicLSTMCellZeroState/zeros_1"}, |
| 46 | {"back_edge_source_array", "rnn/basic_lstm_cell/Mul_2"}, |
| 47 | {"size", "4"}}); |
| 48 | |
| 49 | string message; |
| 50 | ParsedModelFlags result_flags; |
| 51 | |
| 52 | EXPECT_TRUE(ParseModelFlagsFromCommandLineFlags( |
| 53 | &args_count, const_cast<char**>(args), &message, &result_flags)); |
| 54 | EXPECT_EQ(result_flags.input_arrays.value(), expected_input_arrays); |
| 55 | EXPECT_EQ(result_flags.rnn_states.value().elements, expected_rnn_states); |
| 56 | } |
| 57 | |
| 58 | } // namespace |
| 59 | } // namespace toco |
nothing calls this directly
no test coverage detected