| 269 | } |
| 270 | |
| 271 | TfLiteStatus EvalFloat(const TfLiteTensor* input, const TfLiteTensor* bw_input, |
| 272 | const TfLiteTensor* fw_input_weights, |
| 273 | const TfLiteTensor* fw_recurrent_weights, |
| 274 | const TfLiteTensor* fw_bias, |
| 275 | const TfLiteTensor* bw_input_weights, |
| 276 | const TfLiteTensor* bw_recurrent_weights, |
| 277 | const TfLiteTensor* bw_bias, |
| 278 | const TfLiteTensor* aux_input, |
| 279 | const TfLiteTensor* fw_aux_input_weights, |
| 280 | const TfLiteTensor* bw_aux_input_weights, |
| 281 | const TfLiteBidirectionalSequenceRNNParams* params, |
| 282 | TfLiteTensor* fw_hidden_state, TfLiteTensor* fw_output, |
| 283 | TfLiteTensor* bw_hidden_state, TfLiteTensor* bw_output) { |
| 284 | const bool time_major = params->time_major; |
| 285 | const int batch_size = |
| 286 | (time_major) ? input->dims->data[1] : input->dims->data[0]; |
| 287 | const int max_time = |
| 288 | (time_major) ? input->dims->data[0] : input->dims->data[1]; |
| 289 | const int input_size = input->dims->data[2]; |
| 290 | const int aux_input_size = (aux_input) ? aux_input->dims->data[2] : 0; |
| 291 | |
| 292 | const int fw_num_units = fw_input_weights->dims->data[0]; |
| 293 | const float* fw_bias_ptr = fw_bias->data.f; |
| 294 | const float* fw_input_weights_ptr = fw_input_weights->data.f; |
| 295 | const float* fw_recurrent_weights_ptr = fw_recurrent_weights->data.f; |
| 296 | |
| 297 | const int bw_num_units = bw_input_weights->dims->data[0]; |
| 298 | const float* bw_bias_ptr = bw_bias->data.f; |
| 299 | const float* bw_input_weights_ptr = bw_input_weights->data.f; |
| 300 | const float* bw_recurrent_weights_ptr = bw_recurrent_weights->data.f; |
| 301 | |
| 302 | const float* fw_aux_input_weights_ptr = (fw_aux_input_weights != nullptr) |
| 303 | ? fw_aux_input_weights->data.f |
| 304 | : nullptr; |
| 305 | const float* bw_aux_input_weights_ptr = (bw_aux_input_weights != nullptr) |
| 306 | ? bw_aux_input_weights->data.f |
| 307 | : nullptr; |
| 308 | |
| 309 | const int fw_output_step = |
| 310 | params->merge_outputs ? fw_num_units + bw_num_units : fw_num_units; |
| 311 | const int bw_output_step = |
| 312 | params->merge_outputs ? fw_num_units + bw_num_units : bw_num_units; |
| 313 | if (time_major) { |
| 314 | // Forward cell. |
| 315 | float* fw_hidden_state_ptr_batch = fw_hidden_state->data.f; |
| 316 | for (int s = 0; s < max_time; s++) { |
| 317 | const float* input_ptr_batch = |
| 318 | input->data.f + s * input_size * batch_size; |
| 319 | const float* aux_input_ptr_batch = |
| 320 | (aux_input != nullptr) |
| 321 | ? aux_input->data.f + s * input_size * batch_size |
| 322 | : nullptr; |
| 323 | float* output_ptr_batch = |
| 324 | fw_output->data.f + s * fw_output_step * batch_size; |
| 325 | |
| 326 | kernel_utils::RnnBatchStep( |
| 327 | input_ptr_batch, fw_input_weights_ptr, aux_input_ptr_batch, |
| 328 | fw_aux_input_weights_ptr, fw_recurrent_weights_ptr, fw_bias_ptr, |