Skip to content

Commit 9f095f4

Browse files
BalintCsalaBalintCsala
authored andcommitted
PR comments
1 parent e39b4b3 commit 9f095f4

10 files changed

Lines changed: 223 additions & 139 deletions

projects/miopen/src/hip/utilocl.cpp

Lines changed: 59 additions & 42 deletions
Original file line numberDiff line numberDiff line change
@@ -507,33 +507,51 @@ float Im3d2ColGPU(const Handle& handle,
507507

508508
auto&& kernels = handle.GetKernels("miopenIm3d2Col", network_config);
509509

510-
int im_offset_pack = im_offset;
511-
int im_c_pack = im_c;
510+
const uint64_t im_offset_pack = im_offset;
511+
const uint64_t im_c_pack = im_c;
512+
const uint64_t im_d_pack = im_d;
513+
const uint64_t im_h_pack = im_h;
514+
const uint64_t im_w_pack = im_w;
515+
const uint64_t wei_d_pack = wei_d;
516+
const uint64_t wei_h_pack = wei_h;
517+
const uint64_t wei_w_pack = wei_w;
518+
const uint64_t out_d_pack = out_d;
519+
const uint64_t out_h_pack = out_h;
520+
const uint64_t out_w_pack = out_w;
521+
const uint64_t pad_d_pack = pad_d;
522+
const uint64_t pad_h_pack = pad_h;
523+
const uint64_t pad_w_pack = pad_w;
524+
const uint64_t stride_d_pack = stride_d;
525+
const uint64_t stride_h_pack = stride_h;
526+
const uint64_t stride_w_pack = stride_w;
527+
const uint64_t dilation_d_pack = dilation_d;
528+
const uint64_t dilation_h_pack = dilation_h;
529+
const uint64_t dilation_w_pack = dilation_w;
512530

513531
if(!kernels.empty())
514532
{
515533
auto kernel = kernels.front();
516534
kernel(im,
517535
im_offset_pack,
518536
im_c_pack,
519-
im_d,
520-
im_h,
521-
im_w,
522-
wei_d,
523-
wei_h,
524-
wei_w,
525-
out_d,
526-
out_h,
527-
out_w,
528-
pad_d,
529-
pad_h,
530-
pad_w,
531-
stride_d,
532-
stride_h,
533-
stride_w,
534-
dilation_d,
535-
dilation_h,
536-
dilation_w,
537+
im_d_pack,
538+
im_h_pack,
539+
im_w_pack,
540+
wei_d_pack,
541+
wei_h_pack,
542+
wei_w_pack,
543+
out_d_pack,
544+
out_h_pack,
545+
out_w_pack,
546+
pad_d_pack,
547+
pad_h_pack,
548+
pad_w_pack,
549+
stride_d_pack,
550+
stride_h_pack,
551+
stride_w_pack,
552+
dilation_d_pack,
553+
dilation_h_pack,
554+
dilation_w_pack,
537555
col);
538556
}
539557
else
@@ -546,10 +564,9 @@ float Im3d2ColGPU(const Handle& handle,
546564
add_params(" -DLAYOUT_NHWC=" + std::to_string(static_cast<int>(layoutNHWC)));
547565
add_params(" -DGROUPS=" + std::to_string(num_groups));
548566

549-
size_t global_threads = std::min(
550-
256 * static_cast<std::size_t>(out_d * out_h * out_w * im_c * wei_d * wei_h * wei_w) /
551-
8,
552-
static_cast<std::size_t>(256) * 1024);
567+
size_t global_threads = std::min(256 * static_cast<std::size_t>(out_d) * out_h * out_w *
568+
im_c * wei_d * wei_h * wei_w / 8,
569+
static_cast<std::size_t>(256) * 1024);
553570
const size_t local_threads = std::min(global_threads, static_cast<std::size_t>(256));
554571
if(global_threads % local_threads != 0)
555572
{
@@ -563,24 +580,24 @@ float Im3d2ColGPU(const Handle& handle,
563580
im,
564581
im_offset_pack,
565582
im_c_pack,
566-
im_d,
567-
im_h,
568-
im_w,
569-
wei_d,
570-
wei_h,
571-
wei_w,
572-
out_d,
573-
out_h,
574-
out_w,
575-
pad_d,
576-
pad_h,
577-
pad_w,
578-
stride_d,
579-
stride_h,
580-
stride_w,
581-
dilation_d,
582-
dilation_h,
583-
dilation_w,
583+
im_d_pack,
584+
im_h_pack,
585+
im_w_pack,
586+
wei_d_pack,
587+
wei_h_pack,
588+
wei_w_pack,
589+
out_d_pack,
590+
out_h_pack,
591+
out_w_pack,
592+
pad_d_pack,
593+
pad_h_pack,
594+
pad_w_pack,
595+
stride_d_pack,
596+
stride_h_pack,
597+
stride_w_pack,
598+
dilation_d_pack,
599+
dilation_h_pack,
600+
dilation_w_pack,
584601
col);
585602
}
586603

projects/miopen/src/kernels/MIOpenIm2d2Col.cpp

Lines changed: 12 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -156,8 +156,8 @@ extern "C" __global__ void Im2d2Col_v2(const int data_size_off,
156156
data_t* im_off = im + im_offset;
157157

158158
// One column per output pixel: rows = wei_h * wei_w * num_ch
159-
const int patch_size = WEI_H * WEI_W * channels_per_group;
160-
const int col_group_total_size = patch_size * out_h * out_w;
159+
const int patch_size = WEI_H * WEI_W * channels_per_group;
160+
const index_t col_group_total_size = (index_t)patch_size * out_h * out_w;
161161

162162
const int base_oh = tile_h * TILE_OUT_H;
163163
const int base_ow = tile_w * TILE_OUT_W;
@@ -222,7 +222,7 @@ extern "C" __global__ void Im2d2Col_v2(const int data_size_off,
222222
const index_t col_idx = patch_offset +
223223
((index_t)kh * WEI_W + kw) * channels_per_group +
224224
group_relative_channel;
225-
col[(col_group_total_size * current_group) + col_idx] = v;
225+
col[col_group_total_size * current_group + col_idx] = v;
226226
}
227227
}
228228
}
@@ -337,12 +337,12 @@ extern "C" __global__ void Im2d2Col_v2(const int data_size_off,
337337
data_t* col)
338338
{
339339
const int lid = threadIdx.x;
340-
const int grp_id = blockIdx.x; // patch id (one output pixel)
340+
const index_t grp_id = blockIdx.x; // patch id (one output pixel)
341341
const int local_size = blockDim.x;
342342

343343
data_t* im_off = im + im_offset;
344344

345-
const int output_size = out_h * out_w;
345+
const index_t output_size = (index_t)out_h * out_w;
346346
if(grp_id >= output_size)
347347
return;
348348

@@ -365,8 +365,8 @@ extern "C" __global__ void Im2d2Col_v2(const int data_size_off,
365365

366366
const int current_group = c / channels_per_group;
367367
const int channel_in_current_group = c % channels_per_group;
368-
const int group_offset = current_group * (output_size * patch_size_per_group);
369-
const int patch_offset = grp_id * patch_size_per_group;
368+
const index_t group_offset = (index_t)current_group * output_size * patch_size_per_group;
369+
const index_t patch_offset = (index_t)grp_id * patch_size_per_group;
370370

371371
const int src_h = oh * stride_h + kh * dilation_h - pad_h;
372372
const int src_w = ow * stride_w + kw * dilation_w - pad_w;
@@ -375,12 +375,13 @@ extern "C" __global__ void Im2d2Col_v2(const int data_size_off,
375375
data_t v = (data_t)0;
376376
if(ok)
377377
{
378-
const int input_idx = ((src_h * w + src_w) * CHANNELS) + c;
379-
v = im_off[input_idx];
378+
const index_t input_idx = ((index_t)src_h * w + src_w) * CHANNELS + c;
379+
v = im_off[input_idx];
380380
}
381381

382-
const int col_idx = group_offset + patch_offset + (kh * WEI_W + kw) * channels_per_group +
383-
channel_in_current_group;
382+
const index_t col_idx = group_offset + patch_offset +
383+
((index_t)kh * WEI_W + kw) * channels_per_group +
384+
channel_in_current_group;
384385
col[col_idx] = v;
385386
}
386387
}

projects/miopen/src/kernels/MIOpenIm3d2Col.cpp

Lines changed: 74 additions & 68 deletions
Original file line numberDiff line numberDiff line change
@@ -62,70 +62,76 @@ typedef float data_t;
6262

6363
#if(LAYOUT_NHWC == 1)
6464
extern "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
141147
extern "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)

projects/miopen/src/solver/conv/gemm.cpp

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1066,8 +1066,7 @@ bool GemmFwdRest::IsApplicable(const ExecutionContext& context,
10661066
problem.GetWeights().GetType() != miopenInt8)
10671067
return true;
10681068

1069-
// Everything below goes through Im2Col, which indexes x channel-major.
1070-
if(!problem.IsLayoutDefault())
1069+
if(!(problem.IsLayoutDefault() || problem.IsLayoutNHWC()))
10711070
return false;
10721071

10731072
return GetWorkspaceSize(context, problem) > 0;

projects/miopen/src/solver/conv/gemm_wrw.cpp

Lines changed: 2 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -566,18 +566,8 @@ bool GemmWrwUniversal::IsApplicable(const ExecutionContext& context,
566566
if(!(problem.IsLayoutDefault() || problem.IsLayoutNHWC()))
567567
return false;
568568

569-
if(GetWorkspaceSize(context, problem) != 0)
570-
{
571-
if(problem.GetSpatialDims() > 2)
572-
{
573-
return true;
574-
}
575-
else
576-
{
577-
return !GemmWrw1x1_stride1{}.IsApplicable(context, problem);
578-
}
579-
}
580-
return false;
569+
return GetWorkspaceSize(context, problem) != 0 &&
570+
!GemmWrw1x1_stride1{}.IsApplicable(context, problem);
581571
#else
582572
std::ignore = context;
583573
std::ignore = problem;

0 commit comments

Comments
 (0)