Team Ai
Apppublic

aikenml/data_mining

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
correlation.cpp179 linesDownload Raw Back to Correlation_Module
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