diff options
Diffstat (limited to 'tests/validation/fixtures')
-rw-r--r-- | tests/validation/fixtures/ReduceMeanFixture.h | 3 |
1 files changed, 2 insertions, 1 deletions
diff --git a/tests/validation/fixtures/ReduceMeanFixture.h b/tests/validation/fixtures/ReduceMeanFixture.h index 8692213641..769d7f674f 100644 --- a/tests/validation/fixtures/ReduceMeanFixture.h +++ b/tests/validation/fixtures/ReduceMeanFixture.h @@ -119,9 +119,10 @@ protected: if(!keep_dims) { TensorShape output_shape = src_shape; + std::sort(axis.begin(), axis.begin() + axis.num_dimensions()); for(unsigned int i = 0; i < axis.num_dimensions(); ++i) { - output_shape.remove_dimension(axis[i]); + output_shape.remove_dimension(axis[i] - i); } out = reference::reshape_layer(out, output_shape); |