| 74 | } |
| 75 | |
| 76 | TfLiteStatus Prepare(TfLiteContext* context, TfLiteNode* node) { |
| 77 | const auto* params = reinterpret_cast<TfLiteSVDFParams*>(node->builtin_data); |
| 78 | OpData* op_data = reinterpret_cast<OpData*>(node->user_data); |
| 79 | int scratch_tensor_index = op_data->scratch_tensor_index; |
| 80 | |
| 81 | // Check we have all the inputs and outputs we need. |
| 82 | TF_LITE_ENSURE_EQ(context, node->outputs->size, 1); |
| 83 | TF_LITE_ENSURE_EQ(context, node->inputs->size, 5); |
| 84 | op_data->activation_state_tensor_index = |
| 85 | node->inputs->data[kInputActivationStateTensor]; |
| 86 | |
| 87 | const TfLiteTensor* input = GetInput(context, node, kInputTensor); |
| 88 | const TfLiteTensor* weights_feature = |
| 89 | GetInput(context, node, kWeightsFeatureTensor); |
| 90 | const TfLiteTensor* weights_time = |
| 91 | GetInput(context, node, kWeightsTimeTensor); |
| 92 | |
| 93 | TF_LITE_ENSURE_EQ(context, input->type, kTfLiteFloat32); |
| 94 | |
| 95 | // Check all the parameters of tensor match within themselves and match the |
| 96 | // input configuration. |
| 97 | const int rank = params->rank; |
| 98 | const int batch_size = input->dims->data[0]; |
| 99 | const int num_filters = weights_feature->dims->data[0]; |
| 100 | TF_LITE_ENSURE_EQ(context, num_filters % rank, 0); |
| 101 | const int num_units = num_filters / rank; |
| 102 | const int memory_size = weights_time->dims->data[1]; |
| 103 | TF_LITE_ENSURE_EQ(context, input->dims->data[1], |
| 104 | weights_feature->dims->data[1]); |
| 105 | TF_LITE_ENSURE_EQ(context, weights_time->dims->data[0], num_filters); |
| 106 | |
| 107 | const TfLiteTensor* bias = GetOptionalInputTensor(context, node, kBiasTensor); |
| 108 | if (bias) { |
| 109 | TF_LITE_ENSURE_EQ(context, bias->dims->data[0], num_units); |
| 110 | } |
| 111 | |
| 112 | TfLiteTensor* activation_state = |
| 113 | &context->tensors[op_data->activation_state_tensor_index]; |
| 114 | TfLiteTensor* output = GetOutput(context, node, kOutputTensor); |
| 115 | |
| 116 | // Check the shape of input state tensors. |
| 117 | TF_LITE_ENSURE_EQ(context, NumDimensions(activation_state), 2); |
| 118 | TF_LITE_ENSURE_EQ(context, SizeOfDimension(activation_state, 0), batch_size); |
| 119 | TF_LITE_ENSURE_EQ(context, SizeOfDimension(activation_state, 1), |
| 120 | memory_size * num_filters); |
| 121 | |
| 122 | // Resize output. |
| 123 | TfLiteIntArray* output_size_array = TfLiteIntArrayCreate(2); |
| 124 | output_size_array->data[0] = batch_size; |
| 125 | output_size_array->data[1] = num_units; |
| 126 | TF_LITE_ENSURE_OK(context, |
| 127 | context->ResizeTensor(context, output, output_size_array)); |
| 128 | |
| 129 | // The weights are of consistent type, so it suffices to check one. |
| 130 | const bool is_hybrid_op = IsHybridOp(input, weights_feature); |
| 131 | |
| 132 | // Resize scratch. |
| 133 | TfLiteIntArrayFree(node->temporaries); |
nothing calls this directly
no test coverage detected