aikenml/data_mining
0
1#include <torch/extension.h>2using namespace torch;3 4#include <vector>5 6#define WITHIN_BOUNDS(x, y, H, W) (x >= 0 && x < H && y >= 0 && y < W)7 8template <typename scalar_t>9static void correlate_patch(10 TensorAccessor<scalar_t,3> input1,11 TensorAccessor<scalar_t,3> input2,12 scalar_t *dst,13 int kH, int kW,14 int dilationH, int dilationW,15 int u, int v,16 int shiftU, int shiftV){17 const int C = input1.size(0);18 const int iH = input1.size(1);19 const int iW = input1.size(2);20 for (int c=0; c<C; ++c){21 for (int i=0; i<kH; ++i){22 int i1 = u + i * dilationH;23 int i2 = i1 + shiftU;24 if WITHIN_BOUNDS(i1, i2, iH, iH){25 for (int j=0; j<kW; ++j){26 int j1 = v + j * dilationW;27 int j2 = j1 + shiftV;28 if WITHIN_BOUNDS(j1, j2, iW, iW){29 scalar_t v1 = input1[c][i1][j1];30 scalar_t v2 = input2[c][i2][j2];31 *dst += v1 * v2;32 }33 }34 }35 }36 }37}38 39template <typename scalar_t>40static void correlate_patch_grad(41 TensorAccessor<scalar_t,3> input1,42 TensorAccessor<scalar_t,3> gradInput1,43 TensorAccessor<scalar_t,3> input2,44 TensorAccessor<scalar_t,3> gradInput2,45 scalar_t gradOutput,46 int kH, int kW,47 int dilationH, int dilationW,48 int u, int v,49 int shiftU, int shiftV){50 51 const int C = input1.size(0);52 const int iH = input1.size(1);53 const int iW = input1.size(2);54 55 for (int c=0; c<C; ++c){56 for (int i=0; i<kH; ++i){57 int i1 = u + i * dilationH;58 int i2 = i1 + shiftU;59 if WITHIN_BOUNDS(i1, i2, iH, iH){60 for (int j=0; j<kW; ++j){61 int j1 = v + j * dilationW;62 int j2 = j1 + shiftV;63 if WITHIN_BOUNDS(j1, j2, iW, iW){64 scalar_t v1 = input1[c][i1][j1];65 scalar_t v2 = input2[c][i2][j2];66 gradInput2[c][i2][j2] += gradOutput * v1;67 gradInput1[c][i1][j1] += gradOutput * v2;68 }69 }70 }71 }72 }73}74 75torch::Tensor correlation_cpp_forward(76 torch::Tensor input1,77 torch::Tensor input2,78 int kH, int kW,79 int patchH, int patchW,80 int padH, int padW,81 int dilationH, int dilationW,82 int dilation_patchH, int dilation_patchW,83 int dH, int dW) {84 85 const auto batch_size = input1.size(0);86 const auto iH = input1.size(2);87 const auto iW = input1.size(3);88 const int patchRadH = (patchH - 1) / 2;89 const int patchRadW = (patchW - 1) / 2;90 const int dilatedKH = (kH - 1) * dilationH + 1;91 const int dilatedKW = (kW - 1) * dilationW + 1;92 93 const auto oH = (iH + 2 * padH - dilatedKH) / dH + 1;94 const auto oW = (iW + 2 * padW - dilatedKW) / dW + 1;95 auto output = at::zeros({batch_size, patchH, patchW, oH, oW}, input1.options());96 97 int n, ph, pw, h, w;98 #pragma omp parallel for private(n, ph, pw, h, w) collapse(2)99 for (n = 0; n < batch_size; ++n) {100 for(ph = 0; ph < patchH; ++ph){101 for(pw = 0; pw < patchW; ++pw){102 AT_DISPATCH_FLOATING_TYPES(input1.scalar_type(), "correlation_forward_cpp", ([&] {103 auto input1_acc = input1.accessor<scalar_t, 4>();104 auto input2_acc = input2.accessor<scalar_t, 4>();105 auto output_acc = output.accessor<scalar_t, 5>();106 for (h = 0; h < oH; ++h) {107 for (w = 0; w < oW; ++w) {108 correlate_patch(input1_acc[n],109 input2_acc[n],110 &output_acc[n][ph][pw][h][w],111 kH, kW,112 dilationH, dilationW,113 -padH + h * dH,114 -padW + w * dW,115 (ph - patchRadH) * dilation_patchH,116 (pw - patchRadW) * dilation_patchW);117 }118 }119 }));120 }121 }122 }123 return output;124}125 126std::vector<torch::Tensor> correlation_cpp_backward(127 torch::Tensor input1,128 torch::Tensor input2,129 torch::Tensor gradOutput,130 int kH, int kW,131 int patchH, int patchW,132 int padH, int padW,133 int dilationH, int dilationW,134 int dilation_patchH, int dilation_patchW,135 int dH, int dW) {136 137 const int batch_size = input1.size(0);138 const int patchRadH = (patchH - 1) / 2;139 const int patchRadW = (patchW - 1) / 2;140 const int oH = gradOutput.size(3);141 const int oW = gradOutput.size(4);142 143 auto gradInput1 = torch::zeros_like(input1);144 145 auto gradInput2 = torch::zeros_like(input2);146 147 int n, ph, pw, h, w;148 #pragma omp parallel for private(n, ph, pw, h, w)149 for (n = 0; n < batch_size; ++n) {150 AT_DISPATCH_FLOATING_TYPES(input1.scalar_type(), "correlation_backward_cpp", ([&] {151 auto input1_acc = input1.accessor<scalar_t, 4>();152 auto gradInput1_acc = gradInput1.accessor<scalar_t, 4>();153 auto input2_acc = input2.accessor<scalar_t, 4>();154 auto gradInput2_acc = gradInput2.accessor<scalar_t, 4>();155 auto gradOutput_acc = gradOutput.accessor<scalar_t, 5>();156 157 for(ph = 0; ph < patchH; ++ph){158 for(pw = 0; pw < patchW; ++pw){159 for (h = 0; h < oH; ++h) {160 for (w = 0; w < oW; ++w) {161 correlate_patch_grad(input1_acc[n], gradInput1_acc[n],162 input2_acc[n], gradInput2_acc[n],163 gradOutput_acc[n][ph][pw][h][w],164 kH, kW,165 dilationH, dilationW,166 -padH + h * dH,167 -padW + w * dW,168 (ph - patchRadH) * dilation_patchH,169 (pw - patchRadW) * dilation_patchW);170 }171 }172 }173 }174 }));175 }176 177 return {gradInput1, gradInput2};178}179 