Resizes output array based on the input size and resolved axis.
| 93 | |
| 94 | // Resizes output array based on the input size and resolved axis. |
| 95 | TfLiteStatus ResizeOutputTensor(TfLiteContext* context, OpContext* op_context) { |
| 96 | size_t num_axis = NumElements(op_context->axis); |
| 97 | const TfLiteIntArray* input_dims = op_context->input->dims; |
| 98 | int input_num_dims = NumDimensions(op_context->input); |
| 99 | if (input_num_dims == 0) { |
| 100 | return context->ResizeTensor(context, op_context->output, |
| 101 | TfLiteIntArrayCreate(0)); |
| 102 | } |
| 103 | const int* axis = GetTensorData<int>(op_context->axis); |
| 104 | if (op_context->params->keep_dims) { |
| 105 | TfLiteIntArray* output_dims = TfLiteIntArrayCreate(input_num_dims); |
| 106 | for (int idx = 0; idx < input_num_dims; ++idx) { |
| 107 | bool is_axis = false; |
| 108 | for (int axis_idx = 0; axis_idx < num_axis; ++axis_idx) { |
| 109 | if (axis[axis_idx] == idx || axis[axis_idx] + input_num_dims == idx) { |
| 110 | is_axis = true; |
| 111 | break; |
| 112 | } |
| 113 | } |
| 114 | if (is_axis) { |
| 115 | output_dims->data[idx] = 1; |
| 116 | } else { |
| 117 | output_dims->data[idx] = input_dims->data[idx]; |
| 118 | } |
| 119 | } |
| 120 | return context->ResizeTensor(context, op_context->output, output_dims); |
| 121 | } else { |
| 122 | // Calculates size of reducing axis. |
| 123 | int num_reduce_axis = num_axis; |
| 124 | for (int i = 0; i < num_axis; ++i) { |
| 125 | int current = axis[i]; |
| 126 | if (current < 0) { |
| 127 | current += input_num_dims; |
| 128 | } |
| 129 | TF_LITE_ENSURE(context, current >= 0 && current < input_num_dims); |
| 130 | for (int j = 0; j < i; ++j) { |
| 131 | int previous = axis[j]; |
| 132 | if (previous < 0) { |
| 133 | previous += input_num_dims; |
| 134 | } |
| 135 | if (current == previous) { |
| 136 | --num_reduce_axis; |
| 137 | break; |
| 138 | } |
| 139 | } |
| 140 | } |
| 141 | // Determines output dimensions. |
| 142 | TfLiteIntArray* output_dims = |
| 143 | TfLiteIntArrayCreate(input_num_dims - num_reduce_axis); |
| 144 | int num_skip_axis = 0; |
| 145 | for (int idx = 0; idx < input_num_dims; ++idx) { |
| 146 | bool is_axis = false; |
| 147 | for (int axis_idx = 0; axis_idx < num_axis; ++axis_idx) { |
| 148 | if (axis[axis_idx] == idx || axis[axis_idx] + input_num_dims == idx) { |
| 149 | ++num_skip_axis; |
| 150 | is_axis = true; |
| 151 | break; |
| 152 | } |
no test coverage detected