| 2310 | } |
| 2311 | |
| 2312 | void FinishBuildingRNNStates(Model* model) { |
| 2313 | for (const auto& rnn_state : model->flags.rnn_states()) { |
| 2314 | if (!model->HasArray(rnn_state.back_edge_source_array()) || |
| 2315 | !model->HasArray(rnn_state.state_array())) { |
| 2316 | CHECK(model->HasArray(rnn_state.back_edge_source_array())); |
| 2317 | CHECK(model->HasArray(rnn_state.state_array())); |
| 2318 | continue; |
| 2319 | } |
| 2320 | const auto& src_array = model->GetArray(rnn_state.back_edge_source_array()); |
| 2321 | auto& dst_array = model->GetArray(rnn_state.state_array()); |
| 2322 | if (src_array.data_type == ArrayDataType::kNone && |
| 2323 | dst_array.data_type == ArrayDataType::kNone) { |
| 2324 | dst_array.data_type = ArrayDataType::kFloat; |
| 2325 | } |
| 2326 | } |
| 2327 | } |
| 2328 | |
| 2329 | // Returns the array names that match the ArraysExtraInfo's name and |
| 2330 | // name_regexp. The regexp match is for a full match. |
no test coverage detected