GPUCodeForces/S1/ZZZJ_#191/digitization_cuda.py

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)