Skip to content

Commit 69eb495

Browse files
authored
Add overflow checks to getLeadingDims and getTrailingDims
Differential Revision: D103467782 Pull Request resolved: #19270
1 parent fa703b4 commit 69eb495

1 file changed

Lines changed: 15 additions & 2 deletions

File tree

runtime/core/exec_aten/util/tensor_util.h

Lines changed: 15 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@
99
#pragma once
1010

1111
#include <c10/util/irange.h>
12+
#include <c10/util/safe_numerics.h>
1213
#include <algorithm>
1314
#include <array> // std::array
1415
#include <cinttypes> // PRId64
@@ -932,7 +933,13 @@ inline size_t getLeadingDims(
932933
ssize_t(tensor.dim()));
933934
size_t dims = 1;
934935
for (const auto i : c10::irange(dim)) {
935-
dims *= static_cast<size_t>(tensor.size(i));
936+
size_t next_dims;
937+
ET_CHECK_MSG(
938+
!c10::mul_overflows(
939+
dims, static_cast<size_t>(tensor.size(i)), &next_dims),
940+
"Overflow computing leading dims at dimension %zd",
941+
(ssize_t)i);
942+
dims = next_dims;
936943
}
937944
return dims;
938945
}
@@ -949,7 +956,13 @@ inline size_t getTrailingDims(
949956
ssize_t(tensor.dim()));
950957
size_t dims = 1;
951958
for (size_t i = dim + 1; i < static_cast<size_t>(tensor.dim()); ++i) {
952-
dims *= static_cast<size_t>(tensor.size(i));
959+
size_t next_dims;
960+
ET_CHECK_MSG(
961+
!c10::mul_overflows(
962+
dims, static_cast<size_t>(tensor.size(i)), &next_dims),
963+
"Overflow computing trailing dims at dimension %zu",
964+
i);
965+
dims = next_dims;
953966
}
954967
return dims;
955968
}

0 commit comments

Comments
 (0)