forked from ccf-ai-infra/GPUCodeForces
179 lines
6.3 KiB
Python
179 lines
6.3 KiB
Python
# l1loss_cuda.py
|
|
import torch
|
|
from torch.utils.cpp_extension import load_inline
|
|
|
|
class ModelNew(torch.nn.Module):
|
|
|
|
def __init__(self):
|
|
super().__init__()
|
|
self._compile_cuda_kernel()
|
|
|
|
def _compile_cuda_kernel(self):
|
|
cpp_source = """
|
|
#include <torch/extension.h>
|
|
|
|
torch::Tensor l1_forward_cuda(torch::Tensor pred, torch::Tensor target);
|
|
"""
|
|
|
|
cuda_source = """
|
|
#include <cuda_runtime.h>
|
|
#include <device_launch_parameters.h>
|
|
|
|
#define BLOCK_SIZE 256
|
|
#define VEC_SIZE 4
|
|
#define WARP_SIZE 32
|
|
|
|
// Warp-level reduction using shuffle instructions
|
|
__device__ __forceinline__ float warp_reduce_sum(float val) {
|
|
#pragma unroll
|
|
for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {
|
|
val += __shfl_down_sync(0xffffffff, val, offset);
|
|
}
|
|
return val;
|
|
}
|
|
|
|
// Optimized L1 kernel with multiple improvements
|
|
__global__ void l1_optimized_kernel(
|
|
const float* __restrict__ pred,
|
|
const float* __restrict__ target,
|
|
float* __restrict__ output_sum,
|
|
int N_elements
|
|
) {
|
|
float thread_sum = 0.0f;
|
|
|
|
int N_vec = N_elements / VEC_SIZE;
|
|
int grid_stride_vec = gridDim.x * blockDim.x;
|
|
|
|
const float4* __restrict__ pred4 = reinterpret_cast<const float4*>(pred);
|
|
const float4* __restrict__ target4 = reinterpret_cast<const float4*>(target);
|
|
|
|
// Grid-stride loop with vectorized loads
|
|
for (int idx_vec = blockIdx.x * blockDim.x + threadIdx.x;
|
|
idx_vec < N_vec;
|
|
idx_vec += grid_stride_vec)
|
|
{
|
|
// Use read-only cache for better memory performance
|
|
float4 p4 = __ldg(&pred4[idx_vec]);
|
|
float4 t4 = __ldg(&target4[idx_vec]);
|
|
|
|
// Use fabsf() instead of std::abs() - much faster on GPU
|
|
// fabsf is a single instruction, while std::abs may have overhead
|
|
thread_sum += fabsf(p4.x - t4.x);
|
|
thread_sum += fabsf(p4.y - t4.y);
|
|
thread_sum += fabsf(p4.z - t4.z);
|
|
thread_sum += fabsf(p4.w - t4.w);
|
|
}
|
|
|
|
// Warp-level reduction (no shared memory for intra-warp)
|
|
thread_sum = warp_reduce_sum(thread_sum);
|
|
|
|
// Shared memory only for inter-warp reduction
|
|
__shared__ float warp_sums[BLOCK_SIZE / WARP_SIZE];
|
|
|
|
int lane = threadIdx.x % WARP_SIZE;
|
|
int warp_id = threadIdx.x / WARP_SIZE;
|
|
|
|
// First thread in each warp writes to shared memory
|
|
if (lane == 0) {
|
|
warp_sums[warp_id] = thread_sum;
|
|
}
|
|
__syncthreads();
|
|
|
|
// Final reduction by first warp
|
|
if (warp_id == 0) {
|
|
thread_sum = (threadIdx.x < BLOCK_SIZE / WARP_SIZE) ? warp_sums[lane] : 0.0f;
|
|
thread_sum = warp_reduce_sum(thread_sum);
|
|
|
|
if (threadIdx.x == 0) {
|
|
output_sum[blockIdx.x] = thread_sum;
|
|
}
|
|
}
|
|
}
|
|
|
|
// Final reduction kernel - sums up partial results on GPU
|
|
__global__ void final_reduction_kernel(
|
|
const float* __restrict__ partial_sums,
|
|
float* __restrict__ output,
|
|
int n
|
|
) {
|
|
__shared__ float sh_sum[BLOCK_SIZE];
|
|
|
|
float sum = 0.0f;
|
|
|
|
// Grid-stride loop to handle any number of partial sums
|
|
for (int i = threadIdx.x; i < n; i += blockDim.x) {
|
|
sum += partial_sums[i];
|
|
}
|
|
|
|
sh_sum[threadIdx.x] = sum;
|
|
__syncthreads();
|
|
|
|
// Tree reduction in shared memory
|
|
#pragma unroll
|
|
for (int s = BLOCK_SIZE / 2; s > 0; s /= 2) {
|
|
if (threadIdx.x < s) {
|
|
sh_sum[threadIdx.x] += sh_sum[threadIdx.x + s];
|
|
}
|
|
__syncthreads();
|
|
}
|
|
|
|
if (threadIdx.x == 0) {
|
|
output[0] = sh_sum[0];
|
|
}
|
|
}
|
|
|
|
torch::Tensor l1_forward_cuda(torch::Tensor pred, torch::Tensor target) {
|
|
TORCH_CHECK(pred.is_cuda() && target.is_cuda(), "Inputs must be CUDA tensors");
|
|
|
|
pred = pred.contiguous();
|
|
target = target.contiguous();
|
|
|
|
int N_elements = pred.numel();
|
|
TORCH_CHECK(N_elements % VEC_SIZE == 0,
|
|
"Total elements must be divisible by VEC_SIZE (4)");
|
|
|
|
// Adaptive grid size based on GPU architecture
|
|
const int block_size = BLOCK_SIZE;
|
|
int num_sms;
|
|
cudaDeviceGetAttribute(&num_sms, cudaDevAttrMultiProcessorCount, 0);
|
|
|
|
// Heuristic: 4 blocks per SM for good occupancy
|
|
const int grid_size = min(num_sms * 4, (N_elements / VEC_SIZE + block_size - 1) / block_size);
|
|
|
|
// Allocate temporary storage for partial sums
|
|
auto partial_sum = torch::empty({grid_size}, pred.options());
|
|
auto final_result = torch::empty({1}, pred.options());
|
|
|
|
// Launch main reduction kernel
|
|
l1_optimized_kernel<<<grid_size, block_size>>>(
|
|
pred.data_ptr<float>(),
|
|
target.data_ptr<float>(),
|
|
partial_sum.data_ptr<float>(),
|
|
N_elements
|
|
);
|
|
|
|
// Launch final reduction kernel (entirely on GPU)
|
|
final_reduction_kernel<<<1, block_size>>>(
|
|
partial_sum.data_ptr<float>(),
|
|
final_result.data_ptr<float>(),
|
|
grid_size
|
|
);
|
|
|
|
// Compute mean on GPU
|
|
final_result.div_(N_elements);
|
|
|
|
return final_result;
|
|
}
|
|
"""
|
|
|
|
self.l1_op = load_inline(
|
|
name="l1loss_optimized_op",
|
|
cpp_sources=cpp_source,
|
|
cuda_sources=cuda_source,
|
|
functions=["l1_forward_cuda"],
|
|
extra_cuda_cflags=["-O3", "--use_fast_math", "-lineinfo"],
|
|
verbose=True
|
|
)
|
|
|
|
def forward(self, pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
|
|
return self.l1_op.l1_forward_cuda(pred, target) |