Skip to content

Commit 1e9981a

Browse files
committed
Fix int32 overflow in max_pool3d backward kernel
Widen parameters and variables in `max_pool3d_with_indices_backward_template` and `MaxPool3dBackwardKernelFunctor` from `int` to `int64_t` to prevent overflow when computing offsets for large tensors (e.g. shape `[70, 32, 100, 100, 100]`). Fixes test case: `test_pool3d_large_size_int64`
1 parent 0247988 commit 1e9981a

1 file changed

Lines changed: 25 additions & 25 deletions

File tree

src/ATen/native/xpu/sycl/DilatedMaxPool3d.cpp

Lines changed: 25 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -321,17 +321,17 @@ struct MaxPool3dBackwardKernelFunctor {
321321
void operator()(sycl::nd_item<1> item) const {
322322
auto outputIndex = item.get_global_id(0);
323323
if (outputIndex < gradOutputSize_) {
324-
int batch = outputIndex / out_nbatch_stride_;
324+
int64_t batch = outputIndex / out_nbatch_stride_;
325325
if constexpr (channels_last) {
326-
int channel = outputIndex % features_;
326+
int64_t channel = outputIndex % features_;
327327
int64_t index = indicesData_[outputIndex];
328328
int64_t gradIn_offset =
329329
batch * in_nbatch_stride_ + channel + index * features_;
330330
atomicAdd(
331331
(sycl_global_ptr<scalar_t>)&gradInputData_[gradIn_offset],
332332
gradOutputData_[outputIndex]);
333333
} else {
334-
int channel = outputIndex / out_cf_channel_stride_ % features_;
334+
int64_t channel = outputIndex / out_cf_channel_stride_ % features_;
335335
int64_t index = indicesData_[outputIndex];
336336
int64_t gradIn_offset =
337337
batch * in_nbatch_stride_ + channel * in_cf_channel_stride_ + index;
@@ -345,12 +345,12 @@ struct MaxPool3dBackwardKernelFunctor {
345345
scalar_t* gradInputData,
346346
const scalar_t* gradOutputData,
347347
const int64_t* indicesData,
348-
int features,
348+
int64_t features,
349349
int64_t gradOutputSize,
350-
int out_cf_channel_stride,
351-
int in_cf_channel_stride,
352-
int out_nbatch_stride,
353-
int in_nbatch_stride)
350+
int64_t out_cf_channel_stride,
351+
int64_t in_cf_channel_stride,
352+
int64_t out_nbatch_stride,
353+
int64_t in_nbatch_stride)
354354
: gradInputData_(gradInputData),
355355
gradOutputData_(gradOutputData),
356356
indicesData_(indicesData),
@@ -365,33 +365,33 @@ struct MaxPool3dBackwardKernelFunctor {
365365
scalar_t* gradInputData_;
366366
const scalar_t* gradOutputData_;
367367
const int64_t* indicesData_;
368-
int features_;
368+
int64_t features_;
369369
int64_t gradOutputSize_;
370-
int out_cf_channel_stride_;
371-
int in_cf_channel_stride_;
372-
int out_nbatch_stride_;
373-
int in_nbatch_stride_;
370+
int64_t out_cf_channel_stride_;
371+
int64_t in_cf_channel_stride_;
372+
int64_t out_nbatch_stride_;
373+
int64_t in_nbatch_stride_;
374374
};
375375

376376
template <typename scalar_t, bool channels_last>
377377
void max_pool3d_with_indices_backward_template(
378378
scalar_t* gradInputData,
379379
const scalar_t* gradOutputData,
380380
const int64_t* indicesData,
381-
int features,
382-
int itime,
383-
int iheight,
384-
int iwidth,
385-
int obatch,
386-
int otime,
387-
int oheight,
388-
int owidth) {
381+
int64_t features,
382+
int64_t itime,
383+
int64_t iheight,
384+
int64_t iwidth,
385+
int64_t obatch,
386+
int64_t otime,
387+
int64_t oheight,
388+
int64_t owidth) {
389389
int64_t gradOutputSize = obatch * features * otime * oheight * owidth;
390390

391-
auto out_cf_channel_stride = otime * oheight * owidth;
392-
auto in_cf_channel_stride = itime * iheight * iwidth;
393-
auto out_nbatch_stride = features * out_cf_channel_stride;
394-
auto in_nbatch_stride = features * in_cf_channel_stride;
391+
int64_t out_cf_channel_stride = otime * oheight * owidth;
392+
int64_t in_cf_channel_stride = itime * iheight * iwidth;
393+
int64_t out_nbatch_stride = features * out_cf_channel_stride;
394+
int64_t in_nbatch_stride = features * in_cf_channel_stride;
395395
MaxPool3dBackwardKernelFunctor<scalar_t, channels_last> kfn(
396396
gradInputData,
397397
gradOutputData,

0 commit comments

Comments
 (0)