| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| #ifndef TENSORFLOW_LITE_KERNELS_INTERNAL_RUNTIME_SHAPE_H_ |
| #define TENSORFLOW_LITE_KERNELS_INTERNAL_RUNTIME_SHAPE_H_ |
|
|
| namespace tflite { |
|
|
| template <int N> |
| struct Dims { |
| int sizes[N]; |
| int strides[N]; |
| }; |
|
|
| class RuntimeShape { |
| public: |
| RuntimeShape& operator=(RuntimeShape const&) = delete; |
|
|
| |
| |
| |
| static constexpr int kMaxSmallSize = 5; |
|
|
| RuntimeShape() : size_(0) {} |
|
|
| explicit RuntimeShape(int dimensions_count) : size_(dimensions_count) {} |
|
|
| RuntimeShape(int shape_size, int32_t value) : size_(shape_size) { |
| for (int i = 0; i < shape_size; ++i) { |
| SetDim(i, value); |
| } |
| } |
|
|
| RuntimeShape(int dimensions_count, const int32_t* dims_data) |
| : size_(dimensions_count) { |
| ReplaceWith(dimensions_count, dims_data); |
| } |
|
|
| bool operator==(const RuntimeShape& comp) const { |
| return this->size_ == comp.size_ && |
| std::memcmp(DimsData(), comp.DimsData(), size_ * sizeof(int32_t)) == |
| 0; |
| } |
|
|
| ~RuntimeShape() {} |
|
|
| int32_t DimensionsCount() const { return size_; } |
| int32_t Dims(int i) const { |
| TFLITE_DCHECK_GE(i, 0); |
| TFLITE_DCHECK_LT(i, size_); |
| return dims_[i]; |
| } |
| void SetDim(int i, int32_t val) { |
| TFLITE_DCHECK_GE(i, 0); |
| TFLITE_DCHECK_LT(i, size_); |
| dims_[i] = val; |
| } |
|
|
| static RuntimeShape ExtendedShape(int new_shape_size, |
| const RuntimeShape& shape) { |
| return RuntimeShape(new_shape_size, shape, 1); |
| } |
| int32_t* DimsData() { return dims_; } |
| const int32_t* DimsData() const { return dims_; } |
| const int32_t* DimsDataUpTo5D() const { return dims_; } |
|
|
| void ReplaceWith(int dimensions_count, const int32_t* dims_data) { |
| size_ = dimensions_count; |
| int32_t* dst_dims = DimsData(); |
| std::memcpy(dst_dims, dims_data, dimensions_count * sizeof(int32_t)); |
| } |
|
|
| |
| |
| int FlatSize() const { |
| int buffer_size = 1; |
| const int* dims_data = reinterpret_cast<const int*>(DimsData()); |
| for (int i = 0; i < size_; i++) { |
| buffer_size *= dims_data[i]; |
| } |
| return buffer_size; |
| } |
|
|
| private: |
| |
| |
| |
| RuntimeShape(int new_shape_size, const RuntimeShape& shape, int pad_value) |
| : size_(new_shape_size) { |
| |
| |
| TFLITE_CHECK_GE(new_shape_size, shape.DimensionsCount()); |
| const int size_increase = new_shape_size - shape.DimensionsCount(); |
| for (int i = 0; i < size_increase; ++i) { |
| SetDim(i, pad_value); |
| } |
| std::memcpy(DimsData() + size_increase, shape.DimsData(), |
| sizeof(int32_t) * shape.DimensionsCount()); |
| } |
|
|
| int32_t size_; |
| union { |
| int32_t dims_[kMaxSmallSize]; |
| }; |
| }; |
|
|
| |
| |
| |
| |
|
|
| inline int Offset(const RuntimeShape& shape, int i0, int i1, int i2, int i3) { |
| TFLITE_DCHECK_EQ(shape.DimensionsCount(), 4); |
| const int* dims_data = reinterpret_cast<const int*>(shape.DimsData()); |
| TFLITE_DCHECK((dims_data[0] == 0 && i0 == 0) || |
| (i0 >= 0 && i0 < dims_data[0])); |
| TFLITE_DCHECK((dims_data[1] == 0 && i1 == 0) || |
| (i1 >= 0 && i1 < dims_data[1])); |
| TFLITE_DCHECK((dims_data[2] == 0 && i2 == 0) || |
| (i2 >= 0 && i2 < dims_data[2])); |
| TFLITE_DCHECK((dims_data[3] == 0 && i3 == 0) || |
| (i3 >= 0 && i3 < dims_data[3])); |
| return ((i0 * dims_data[1] + i1) * dims_data[2] + i2) * dims_data[3] + i3; |
| } |
|
|
| inline int Offset(const RuntimeShape& shape, int i0, int i1, int i2, int i3, |
| int i4) { |
| TFLITE_DCHECK_EQ(shape.DimensionsCount(), 5); |
| const int* dims_data = reinterpret_cast<const int*>(shape.DimsData()); |
| TFLITE_DCHECK((dims_data[0] == 0 && i0 == 0) || |
| (i0 >= 0 && i0 < dims_data[0])); |
| TFLITE_DCHECK((dims_data[1] == 0 && i1 == 0) || |
| (i1 >= 0 && i1 < dims_data[1])); |
| TFLITE_DCHECK((dims_data[2] == 0 && i2 == 0) || |
| (i2 >= 0 && i2 < dims_data[2])); |
| TFLITE_DCHECK((dims_data[3] == 0 && i3 == 0) || |
| (i3 >= 0 && i3 < dims_data[3])); |
| TFLITE_DCHECK((dims_data[4] == 0 && i4 == 0) || |
| (i4 >= 0 && i4 < dims_data[4])); |
| return (((i0 * dims_data[1] + i1) * dims_data[2] + i2) * dims_data[3] + i3) * |
| dims_data[4] + |
| i4; |
| } |
|
|
| } |
|
|
| #endif |
|
|