| 57 | } |
| 58 | |
| 59 | TensorShape ExtendShape(const TensorShape& inputShape, |
| 60 | unsigned int newNumDimensions) |
| 61 | { |
| 62 | if (inputShape.GetNumDimensions() >= newNumDimensions) |
| 63 | { |
| 64 | return inputShape; |
| 65 | } |
| 66 | |
| 67 | std::vector<unsigned int> newSizes(newNumDimensions, 0); |
| 68 | |
| 69 | unsigned int diff = newNumDimensions - inputShape.GetNumDimensions(); |
| 70 | |
| 71 | for (unsigned int i = 0; i < diff; i++) |
| 72 | { |
| 73 | newSizes[i] = 1; |
| 74 | } |
| 75 | |
| 76 | for (unsigned int i = diff; i < newNumDimensions; i++) |
| 77 | { |
| 78 | newSizes[i] = inputShape[i - diff]; |
| 79 | } |
| 80 | |
| 81 | return TensorShape(newNumDimensions, newSizes.data()); |
| 82 | } |
| 83 | |
| 84 | } // Anonymous namespace |
| 85 |
no test coverage detected