31 lines
708 B
Plaintext
31 lines
708 B
Plaintext
#include <ATen/ATen.h>
|
|
|
|
__global__ void my_cuda_kernel(float* input, float* output, int size) {
|
|
int idx = blockIdx.x * blockDim.x + threadIdx.x;
|
|
if (idx < size) {
|
|
output[idx] = 2 * input[idx];
|
|
}
|
|
}
|
|
|
|
at::Tensor my_cuda_function(at::Tensor input) {
|
|
int size = input.numel();
|
|
|
|
// Allocate output tensor on the GPU
|
|
at::Tensor output = at::empty_like(input);
|
|
|
|
// Launch the CUDA kernel
|
|
dim3 blockDim(256);
|
|
dim3 gridDim((size + blockDim.x - 1) / blockDim.x);
|
|
|
|
my_cuda_kernel<<<gridDim, blockDim>>>(
|
|
input.data<float>(),
|
|
output.data<float>(),
|
|
size
|
|
);
|
|
|
|
// Wait for the kernel to finish
|
|
cudaDeviceSynchronize();
|
|
|
|
return output;
|
|
}
|