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

Function Prepare

tensorflow/lite/kernels/svdf.cc:76–193  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

74}
75
76TfLiteStatus 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);

Callers

nothing calls this directly

Calls 13

GetInputFunction · 0.85
GetOptionalInputTensorFunction · 0.85
GetOutputFunction · 0.85
NumDimensionsFunction · 0.85
SizeOfDimensionFunction · 0.85
TfLiteIntArrayCreateFunction · 0.85
IsHybridOpFunction · 0.85
TfLiteIntArrayFreeFunction · 0.85
GetTemporaryFunction · 0.85
TfLiteIntArrayEqualFunction · 0.85
TfLiteIntArrayCopyFunction · 0.85

Tested by

no test coverage detected