|
38 | 38 | #error "MIOPEN_USE_64BIT_INDEX must be defined" |
39 | 39 | #endif |
40 | 40 |
|
| 41 | +#ifndef LAYOUT_NHWC |
| 42 | +#define LAYOUT_NHWC 0 |
| 43 | +#endif |
| 44 | + |
| 45 | +#if(LAYOUT_NHWC == 1) |
| 46 | +extern "C" __global__ void Col2Im3dU(FLOAT* col, |
| 47 | + const unsigned int col_d, |
| 48 | + const unsigned int col_h, |
| 49 | + const unsigned int col_w, |
| 50 | + const unsigned int wei_d, |
| 51 | + const unsigned int wei_h, |
| 52 | + const unsigned int wei_w, |
| 53 | + const unsigned int pad_d, |
| 54 | + const unsigned int pad_h, |
| 55 | + const unsigned int pad_w, |
| 56 | + const unsigned int stride_d, |
| 57 | + const unsigned int stride_h, |
| 58 | + const unsigned int stride_w, |
| 59 | + const unsigned int dilation_d, |
| 60 | + const unsigned int dilation_h, |
| 61 | + const unsigned int dilation_w, |
| 62 | + const unsigned int channels, |
| 63 | + const unsigned int depth, |
| 64 | + const unsigned int height, |
| 65 | + const unsigned int width, |
| 66 | + FLOAT* im, |
| 67 | + const uint64_t im_offset) |
| 68 | +{ |
| 69 | + const unsigned int num_groups = GROUPS; |
| 70 | + const unsigned int channels_per_group = channels / num_groups; |
| 71 | + FLOAT* im_off = im + im_offset; |
| 72 | + unsigned int gid = blockIdx.x * blockDim.x + threadIdx.x; |
| 73 | + unsigned int global_size = channels * depth * height * width; |
| 74 | + if(gid >= global_size) |
| 75 | + return; |
| 76 | + |
| 77 | + unsigned int im_ch = gid % channels; |
| 78 | + unsigned int group_id = im_ch / channels_per_group; |
| 79 | + unsigned int ch_in_group = im_ch % channels_per_group; |
| 80 | + |
| 81 | + unsigned int itmp = gid / channels; |
| 82 | + unsigned int im_w = itmp % width; |
| 83 | + itmp = itmp / width; |
| 84 | + unsigned int im_h = itmp % height; |
| 85 | + unsigned int im_d = itmp / height; |
| 86 | + |
| 87 | + im_d += pad_d; |
| 88 | + im_h += pad_h; |
| 89 | + im_w += pad_w; |
| 90 | + |
| 91 | + unsigned int start_d = (im_d < dilation_d * (wei_d - 1) + 1) |
| 92 | + ? 0 |
| 93 | + : (im_d - (dilation_d * (wei_d - 1) + 1)) / stride_d + 1; |
| 94 | + unsigned int end_d = min(col_d, im_d / stride_d + 1); |
| 95 | + |
| 96 | + unsigned int start_h = (im_h < dilation_h * (wei_h - 1) + 1) |
| 97 | + ? 0 |
| 98 | + : (im_h - (dilation_h * (wei_h - 1) + 1)) / stride_h + 1; |
| 99 | + unsigned int end_h = min(col_h, im_h / stride_h + 1); |
| 100 | + |
| 101 | + unsigned int start_w = (im_w < dilation_w * (wei_w - 1) + 1) |
| 102 | + ? 0 |
| 103 | + : (im_w - (dilation_w * (wei_w - 1) + 1)) / stride_w + 1; |
| 104 | + unsigned int end_w = min(col_w, im_w / stride_w + 1); |
| 105 | + |
| 106 | + uint64_t inner_size = wei_d * wei_h * wei_w * channels_per_group; |
| 107 | + uint64_t col_group_size = col_d * col_h * col_w * inner_size; |
| 108 | + |
| 109 | + FLOAT_ACCUM tmp = (FLOAT_ACCUM)0; |
| 110 | + |
| 111 | + for(unsigned int cz = start_d; cz < end_d; cz++) |
| 112 | + { |
| 113 | + for(unsigned int cy = start_h; cy < end_h; cy++) |
| 114 | + { |
| 115 | + for(unsigned int cx = start_w; cx < end_w; cx++) |
| 116 | + { |
| 117 | + if((im_d - cz * stride_d) % dilation_d == 0 && |
| 118 | + (im_h - cy * stride_h) % dilation_h == 0 && |
| 119 | + (im_w - cx * stride_w) % dilation_w == 0) |
| 120 | + { |
| 121 | + unsigned int z = (im_d - cz * stride_d) / dilation_d; |
| 122 | + unsigned int y = (im_h - cy * stride_h) / dilation_h; |
| 123 | + unsigned int x = (im_w - cx * stride_w) / dilation_w; |
| 124 | + |
| 125 | +#if MIOPEN_USE_64BIT_INDEX |
| 126 | + uint64_t col_off = |
| 127 | + group_id * col_group_size + |
| 128 | + ((((uint64_t)cz * col_h + cy) * col_w + cx) * inner_size) + |
| 129 | + (((uint64_t)z * wei_h + y) * wei_w + x) * channels_per_group + ch_in_group; |
| 130 | + |
| 131 | +#else |
| 132 | + uint32_t col_off = group_id * col_group_size + |
| 133 | + (((cz * col_h + cy) * col_w + cx) * inner_size) + |
| 134 | + ((z * wei_h + y) * wei_w + x) * channels_per_group + |
| 135 | + ch_in_group; |
| 136 | +#endif |
| 137 | + |
| 138 | + tmp += CVT_FLOAT2ACCUM(col[col_off]); |
| 139 | + } |
| 140 | + } |
| 141 | + } |
| 142 | + } |
| 143 | +#if ACCUMULATOR_NEEDS_CONVERSION |
| 144 | + im_off[gid] = tmp > CVT_FLOAT2ACCUM(MAX_VAL) ? MAX_VAL : CVT_ACCUM2FLOAT(tmp); |
| 145 | +#else |
| 146 | + im_off[gid] = tmp; |
| 147 | +#endif |
| 148 | +} |
| 149 | + |
| 150 | +#else |
41 | 151 | extern "C" __global__ void Col2Im3dU(FLOAT* col, |
42 | 152 | const unsigned int col_d, |
43 | 153 | const unsigned int col_h, |
@@ -247,3 +357,4 @@ extern "C" __global__ void Col2Im3dUBatched(FLOAT* col, |
247 | 357 | im_off[localid] = tmp; |
248 | 358 | #endif |
249 | 359 | } |
| 360 | +#endif // #if (LAYOUT_NHWC == 1) |
0 commit comments