@@ -62,70 +62,76 @@ typedef float data_t;
6262
6363#if (LAYOUT_NHWC == 1)
6464extern " C" __global__ void Im3d2Col (data_t * const __restrict im,
65- const unsigned im_offset,
66- const unsigned im_c_size,
67- const unsigned im_d_size,
68- const unsigned im_h_size,
69- const unsigned im_w_size,
70- const unsigned wei_d_size,
71- const unsigned wei_h_size,
72- const unsigned wei_w_size,
73- const unsigned out_d_size,
74- const unsigned out_h_size,
75- const unsigned out_w_size,
76- const unsigned pad_d_size,
77- const unsigned pad_h_size,
78- const unsigned pad_w_size,
79- const unsigned stride_d_size,
80- const unsigned stride_h_size,
81- const unsigned stride_w_size,
82- const unsigned dilation_d_size,
83- const unsigned dilation_h_size,
84- const unsigned dilation_w_size,
65+ const uint64_t im_offset,
66+ const uint64_t im_c_size,
67+ const uint64_t im_d_size,
68+ const uint64_t im_h_size,
69+ const uint64_t im_w_size,
70+ const uint64_t wei_d_size,
71+ const uint64_t wei_h_size,
72+ const uint64_t wei_w_size,
73+ const uint64_t out_d_size,
74+ const uint64_t out_h_size,
75+ const uint64_t out_w_size,
76+ const uint64_t pad_d_size,
77+ const uint64_t pad_h_size,
78+ const uint64_t pad_w_size,
79+ const uint64_t stride_d_size,
80+ const uint64_t stride_h_size,
81+ const uint64_t stride_w_size,
82+ const uint64_t dilation_d_size,
83+ const uint64_t dilation_h_size,
84+ const uint64_t dilation_w_size,
8585 data_t * __restrict col)
8686{
87- const int num_groups = GROUPS ;
88- unsigned channels_per_group = im_c_size / num_groups;
89- unsigned inner_size = wei_d_size * wei_h_size * wei_w_size * channels_per_group;
90- unsigned col_group_size = out_d_size * out_h_size * out_w_size * inner_size;
91- unsigned col_size = col_group_size * num_groups;
87+ const uint64_t num_groups = GROUPS ;
88+ const uint64_t channels_per_group = im_c_size / num_groups;
89+ const uint64_t inner_size = ( uint64_t ) wei_d_size * wei_h_size * wei_w_size * channels_per_group;
90+ const uint64_t col_group_size = ( uint64_t ) out_d_size * out_h_size * out_w_size * inner_size;
91+ const uint64_t col_size = col_group_size * num_groups;
9292
93- unsigned int gtid = blockIdx.x * blockDim.x + threadIdx.x ;
94- unsigned int global_size = blockDim.x * gridDim.x ;
93+ const uint64_t gtid = ( uint64_t ) blockIdx.x * blockDim.x + threadIdx.x ;
94+ const uint64_t global_size = ( uint64_t ) blockDim.x * gridDim.x ;
9595
96- for (unsigned tid = gtid; tid < col_size; tid += global_size)
96+ for (uint64_t tid = gtid; tid < col_size; tid += global_size)
9797 {
98- unsigned group_id = tid / col_group_size;
99- unsigned tid_in_group = tid - group_id * col_group_size;
98+ const uint64_t group_id = tid / col_group_size;
99+ const uint64_t tid_in_group = tid - group_id * col_group_size;
100100
101101 // "col" matrix row and colome id
102- unsigned col_i = tid_in_group / inner_size;
103- unsigned col_j = tid_in_group - col_i * inner_size;
102+ const uint64_t col_i = tid_in_group / inner_size;
103+ const uint64_t col_j = tid_in_group - col_i * inner_size;
104104
105105 // output tensor out_d, out_h, out_w id
106- unsigned out_d = col_i / (out_h_size * out_w_size);
107- unsigned tmp = col_i - out_d * (out_h_size * out_w_size);
108- unsigned out_h = tmp / out_w_size;
109- unsigned out_w = tmp - out_h * out_w_size;
106+ const uint64_t out_hw = (uint64_t )out_h_size * out_w_size;
107+ const uint64_t out_d = col_i / out_hw;
108+ uint64_t tmp = col_i - out_d * out_hw;
109+ const uint64_t out_h = tmp / out_w_size;
110+ const uint64_t out_w = tmp - out_h * out_w_size;
110111
111112 // weight tensor wei_d, wei_h, wei_w, wei_c
112- unsigned wei_d = col_j / (wei_h_size * wei_w_size * channels_per_group);
113- tmp = col_j - wei_d * (wei_h_size * wei_w_size * channels_per_group);
114- unsigned wei_h = tmp / (wei_w_size * channels_per_group);
113+ const uint64_t wei_hwc = (uint64_t )wei_h_size * wei_w_size * channels_per_group;
114+ const uint64_t wei_d = col_j / wei_hwc;
115+ tmp = col_j - wei_d * wei_hwc;
116+ const uint64_t wei_wc = (uint64_t )wei_w_size * channels_per_group;
117+ const uint64_t wei_h = tmp / wei_wc;
115118 tmp -= wei_h * (wei_w_size * channels_per_group);
116- unsigned wei_w = tmp / channels_per_group;
117- unsigned wei_c_in_group = tmp - wei_w * channels_per_group;
119+ const uint64_t wei_w = tmp / channels_per_group;
120+ const uint64_t wei_c_in_group = tmp - wei_w * channels_per_group;
118121
119- unsigned wei_c = wei_c_in_group + group_id * channels_per_group;
122+ const uint64_t wei_c = wei_c_in_group + group_id * channels_per_group;
120123
121124 // input tensor im_d, im_h, im_w id
122- int im_d = (int )(stride_d_size * out_d + dilation_d_size * wei_d) - (int )(pad_d_size);
123- int im_h = (int )(stride_h_size * out_h + dilation_h_size * wei_h) - (int )(pad_h_size);
124- int im_w = (int )(stride_w_size * out_w + dilation_w_size * wei_w) - (int )(pad_w_size);
125+ const int64_t im_d = (int64_t )stride_d_size * (int64_t )out_d +
126+ (int64_t )dilation_d_size * (int64_t )wei_d - (int64_t )pad_d_size;
127+ const int64_t im_h = (int64_t )stride_h_size * (int64_t )out_h +
128+ (int64_t )dilation_h_size * (int64_t )wei_h - (int64_t )pad_h_size;
129+ const int64_t im_w = (int64_t )stride_w_size * (int64_t )out_w +
130+ (int64_t )dilation_w_size * (int64_t )wei_w - (int64_t )pad_w_size;
125131
126- uint64_t im_idx = im_offset + (uint64_t )im_d * ( im_h_size * im_w_size * im_c_size) +
127- ( uint64_t )im_h * (im_w_size * im_c_size) + (uint64_t )im_w * im_c_size +
128- wei_c;
132+ const uint64_t im_idx = im_offset + (uint64_t )im_d * im_h_size * im_w_size * im_c_size +
133+ (uint64_t )im_h * im_w_size * im_c_size +
134+ ( uint64_t )im_w * im_c_size + wei_c;
129135
130136 // NdHWC Memory Access
131137 data_t value = (im_d >= 0 && im_d < im_d_size && im_h >= 0 && im_h < im_h_size &&
@@ -139,26 +145,26 @@ extern "C" __global__ void Im3d2Col(data_t* const __restrict im,
139145
140146#else
141147extern " C" __global__ void Im3d2Col (data_t * const __restrict im,
142- const unsigned im_offset,
143- const unsigned im_c_size,
144- const unsigned im_d_size,
145- const unsigned im_h_size,
146- const unsigned im_w_size,
147- const unsigned wei_d_size,
148- const unsigned wei_h_size,
149- const unsigned wei_w_size,
150- const unsigned out_d_size,
151- const unsigned out_h_size,
152- const unsigned out_w_size,
153- const unsigned pad_d_size,
154- const unsigned pad_h_size,
155- const unsigned pad_w_size,
156- const unsigned stride_d_size,
157- const unsigned stride_h_size,
158- const unsigned stride_w_size,
159- const unsigned dilation_d_size,
160- const unsigned dilation_h_size,
161- const unsigned dilation_w_size,
148+ const uint64_t im_offset,
149+ const uint64_t im_c_size,
150+ const uint64_t im_d_size,
151+ const uint64_t im_h_size,
152+ const uint64_t im_w_size,
153+ const uint64_t wei_d_size,
154+ const uint64_t wei_h_size,
155+ const uint64_t wei_w_size,
156+ const uint64_t out_d_size,
157+ const uint64_t out_h_size,
158+ const uint64_t out_w_size,
159+ const uint64_t pad_d_size,
160+ const uint64_t pad_h_size,
161+ const uint64_t pad_w_size,
162+ const uint64_t stride_d_size,
163+ const uint64_t stride_h_size,
164+ const uint64_t stride_w_size,
165+ const uint64_t dilation_d_size,
166+ const uint64_t dilation_h_size,
167+ const uint64_t dilation_w_size,
162168 data_t * __restrict col)
163169{
164170 // Use size_t to prevent overflow for large tensors (>4GB elements)
0 commit comments