aikenml/data_mining
0
1#include <torch/extension.h>2#include <c10/cuda/CUDAGuard.h>3#include <vector>4#include <iostream>5 6// declarations7 8torch::Tensor correlation_cpp_forward(9 torch::Tensor input1,10 torch::Tensor input2,11 int kH, int kW,12 int patchH, int patchW,13 int padH, int padW,14 int dilationH, int dilationW,15 int dilation_patchH, int dilation_patchW,16 int dH, int dW);17 18std::vector<torch::Tensor> correlation_cpp_backward(19 torch::Tensor grad_output,20 torch::Tensor input1,21 torch::Tensor input2,22 int kH, int kW,23 int patchH, int patchW,24 int padH, int padW,25 int dilationH, int dilationW,26 int dilation_patchH, int dilation_patchW,27 int dH, int dW);28 29#ifdef USE_CUDA30 31#define CHECK_CUDA(x) TORCH_CHECK(x.device().is_cuda(), #x, " must be a CUDA tensor")32#define CHECK_CONTIGUOUS(x) TORCH_CHECK(x.is_contiguous(), #x, " must be contiguous")33#define CHECK_INPUT(x) CHECK_CUDA(x); CHECK_CONTIGUOUS(x)34#define CHECK_SAME_DEVICE(x, y) TORCH_CHECK(x.device() == y.device(), #x " is not on same device as " #y)35 36torch::Tensor correlation_cuda_forward(37 torch::Tensor input1,38 torch::Tensor input2,39 int kH, int kW,40 int patchH, int patchW,41 int padH, int padW,42 int dilationH, int dilationW,43 int dilation_patchH, int dilation_patchW,44 int dH, int dW);45 46std::vector<torch::Tensor> correlation_cuda_backward(47 torch::Tensor grad_output,48 torch::Tensor input1,49 torch::Tensor input2,50 int kH, int kW,51 int patchH, int patchW,52 int padH, int padW,53 int dilationH, int dilationW,54 int dilation_patchH, int dilation_patchW,55 int dH, int dW);56 57// C++ interface58 59torch::Tensor correlation_sample_forward(60 torch::Tensor input1,61 torch::Tensor input2,62 int kH, int kW,63 int patchH, int patchW,64 int padH, int padW,65 int dilationH, int dilationW,66 int dilation_patchH, int dilation_patchW,67 int dH, int dW) {68 if (input1.device().is_cuda()){69 CHECK_INPUT(input1);70 CHECK_INPUT(input2);71 72 // set device of input1 as default CUDA device73 // https://pytorch.org/cppdocs/api/structc10_1_1cuda_1_1_optional_c_u_d_a_guard.html74 const at::cuda::OptionalCUDAGuard guard_input1(device_of(input1));75 CHECK_SAME_DEVICE(input1, input2);76 77 return correlation_cuda_forward(input1, input2, kH, kW, patchH, patchW,78 padH, padW, dilationH, dilationW,79 dilation_patchH, dilation_patchW,80 dH, dW);81 }else{82 return correlation_cpp_forward(input1, input2, kH, kW, patchH, patchW,83 padH, padW, dilationH, dilationW,84 dilation_patchH, dilation_patchW,85 dH, dW);86 }87}88 89std::vector<torch::Tensor> correlation_sample_backward(90 torch::Tensor input1,91 torch::Tensor input2,92 torch::Tensor grad_output,93 int kH, int kW,94 int patchH, int patchW,95 int padH, int padW,96 int dilationH, int dilationW,97 int dilation_patchH, int dilation_patchW,98 int dH, int dW) {99 100 if(grad_output.device().is_cuda()){101 CHECK_INPUT(input1);102 CHECK_INPUT(input2);103 104 // set device of input1 as default CUDA device105 const at::cuda::OptionalCUDAGuard guard_input1(device_of(input1));106 CHECK_SAME_DEVICE(input1, input2);107 CHECK_SAME_DEVICE(input1, grad_output);108 109 return correlation_cuda_backward(input1, input2, grad_output,110 kH, kW, patchH, patchW,111 padH, padW,112 dilationH, dilationW,113 dilation_patchH, dilation_patchW,114 dH, dW);115 }else{116 return correlation_cpp_backward(117 input1, input2, grad_output,118 kH, kW, patchH, patchW,119 padH, padW,120 dilationH, dilationW,121 dilation_patchH, dilation_patchW,122 dH, dW);123 }124}125 126PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {127 m.def("forward", &correlation_sample_forward, "Spatial Correlation Sampler Forward");128 m.def("backward", &correlation_sample_backward, "Spatial Correlation Sampler backward");129}130 131#else132 133PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {134 m.def("forward", &correlation_cpp_forward, "Spatial Correlation Sampler Forward");135 m.def("backward", &correlation_cpp_backward, "Spatial Correlation Sampler backward");136}137 138#endif139 