forked from ccf-ai-infra/GPUCodeForces
102 lines
2.6 KiB
Python
102 lines
2.6 KiB
Python
import torch
|
|
import torch.nn as nn
|
|
from torch.utils.cpp_extension import load_inline
|
|
|
|
# C++ 接口
|
|
cpp_src = "torch::Tensor digitization_cuda(torch::Tensor input, torch::Tensor boundaries);"
|
|
|
|
# CUDA Kernel
|
|
cuda_src = """
|
|
#include <cuda_runtime.h>
|
|
#include <cstdint>
|
|
|
|
extern __shared__ float s_bounds[];
|
|
|
|
__device__ __forceinline__ int binary_search_device(int n, float val) {
|
|
int left = 0;
|
|
int right = n;
|
|
|
|
while (left < right) {
|
|
int mid = left + (right - left) / 2;
|
|
if (s_bounds[mid] < val) {
|
|
left = mid + 1;
|
|
} else {
|
|
right = mid;
|
|
}
|
|
}
|
|
return left;
|
|
}
|
|
|
|
__global__ void digitization_kernel(
|
|
const float* __restrict__ input,
|
|
const float* __restrict__ boundaries,
|
|
int64_t* __restrict__ output,
|
|
int num_elements,
|
|
int num_bounds
|
|
) {
|
|
|
|
for (int i = threadIdx.x; i < num_bounds; i += blockDim.x) {
|
|
s_bounds[i] = boundaries[i];
|
|
}
|
|
|
|
__syncthreads();
|
|
|
|
int idx = blockIdx.x * blockDim.x + threadIdx.x;
|
|
int stride = blockDim.x * gridDim.x;
|
|
|
|
for (int i = idx; i < num_elements; i += stride) {
|
|
float val = input[i];
|
|
|
|
int bin_idx = binary_search_device(num_bounds, val);
|
|
|
|
output[i] = (int64_t)bin_idx;
|
|
}
|
|
}
|
|
|
|
torch::Tensor digitization_cuda(torch::Tensor input, torch::Tensor boundaries) {
|
|
int num_elements = input.numel();
|
|
int num_bounds = boundaries.numel();
|
|
|
|
input = input.contiguous();
|
|
boundaries = boundaries.contiguous();
|
|
|
|
|
|
auto output = torch::empty_like(input, torch::kLong);
|
|
|
|
const int block_size = 256;
|
|
int grid_size = (num_elements + block_size - 1) / block_size;
|
|
if (grid_size > 4096) grid_size = 4096;
|
|
|
|
int shared_mem_bytes = num_bounds * sizeof(float);
|
|
|
|
digitization_kernel<<<grid_size, block_size, shared_mem_bytes>>>(
|
|
input.data_ptr<float>(),
|
|
boundaries.data_ptr<float>(),
|
|
output.data_ptr<int64_t>(),
|
|
num_elements,
|
|
num_bounds
|
|
);
|
|
|
|
return output;
|
|
}
|
|
"""
|
|
|
|
class ModelNew(nn.Module):
|
|
def __init__(self, num_bins=128):
|
|
super().__init__()
|
|
self.boundaries = nn.Parameter(
|
|
torch.linspace(0, 100, steps=num_bins),
|
|
requires_grad=False
|
|
)
|
|
|
|
self.module = load_inline(
|
|
name="digitization_opt_v2_fix",
|
|
cpp_sources=cpp_src,
|
|
cuda_sources=cuda_src,
|
|
functions=["digitization_cuda"],
|
|
verbose=False,
|
|
extra_cuda_cflags=["-O3"]
|
|
)
|
|
|
|
def forward(self, x):
|
|
return self.module.digitization_cuda(x, self.boundaries) |