MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / EvalFloat

Function EvalFloat

tensorflow/lite/kernels/bidirectional_sequence_rnn.cc:271–403  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

269}
270
271TfLiteStatus 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,

Callers 1

EvalFunction · 0.70

Calls 1

RnnBatchStepFunction · 0.85

Tested by

no test coverage detected