GPUCodeForces/S1/uucoco_#21/SoftClip_cuda.py

141 lines
5.3 KiB
Python

import torch
import torch.nn as nn
from torch.utils.cpp_extension import load_inline
class ModelNew(nn.Module):
def __init__(self, min_val=-1.0, max_val=1.0, beta=1.0):
super().__init__()
self.min_val = min_val
self.max_val = max_val
self.beta = beta
self._compile_cuda_kernel()
def _compile_cuda_kernel(self):
cpp_source = """
#include <torch/extension.h>
torch::Tensor softclip_cuda(torch::Tensor x, float min_val, float max_val, float beta);
"""
cuda_source = """
#include <cuda_runtime.h>
__device__ __forceinline__ float softplus_f(float x, float beta, float inv_beta) {
float bx = x * beta;
if (bx > 20.0f) return x;
return log1pf(expf(bx)) * inv_beta;
}
__global__ void softclip_kernel_ilp(
const float* __restrict__ x,
float* __restrict__ y,
int n_vec,
float min_val,
float max_val,
float beta,
float inv_beta)
{
int idx = blockIdx.x * blockDim.x + threadIdx.x;
int stride = gridDim.x * blockDim.x;
const float4* x_vec = reinterpret_cast<const float4*>(x);
float4* y_vec = reinterpret_cast<float4*>(y);
int i = idx;
// ILP Loop Unrolling (2x float4)
for (; i < n_vec - 1; i += stride) {
float4 v1 = x_vec[i];
float4 v2 = x_vec[i+1];
float4 o1, o2;
o1.x = v1.x + softplus_f(min_val - v1.x, beta, inv_beta) - softplus_f(v1.x - max_val, beta, inv_beta);
o1.y = v1.y + softplus_f(min_val - v1.y, beta, inv_beta) - softplus_f(v1.y - max_val, beta, inv_beta);
o1.z = v1.z + softplus_f(min_val - v1.z, beta, inv_beta) - softplus_f(v1.z - max_val, beta, inv_beta);
o1.w = v1.w + softplus_f(min_val - v1.w, beta, inv_beta) - softplus_f(v1.w - max_val, beta, inv_beta);
o2.x = v2.x + softplus_f(min_val - v2.x, beta, inv_beta) - softplus_f(v2.x - max_val, beta, inv_beta);
o2.y = v2.y + softplus_f(min_val - v2.y, beta, inv_beta) - softplus_f(v2.y - max_val, beta, inv_beta);
o2.z = v2.z + softplus_f(min_val - v2.z, beta, inv_beta) - softplus_f(v2.z - max_val, beta, inv_beta);
o2.w = v2.w + softplus_f(min_val - v2.w, beta, inv_beta) - softplus_f(v2.w - max_val, beta, inv_beta);
y_vec[i] = o1;
y_vec[i+1] = o2;
i++;
}
for (; i < n_vec; i += stride) {
float4 v = x_vec[i];
float4 o;
o.x = v.x + softplus_f(min_val - v.x, beta, inv_beta) - softplus_f(v.x - max_val, beta, inv_beta);
o.y = v.y + softplus_f(min_val - v.y, beta, inv_beta) - softplus_f(v.y - max_val, beta, inv_beta);
o.z = v.z + softplus_f(min_val - v.z, beta, inv_beta) - softplus_f(v.z - max_val, beta, inv_beta);
o.w = v.w + softplus_f(min_val - v.w, beta, inv_beta) - softplus_f(v.w - max_val, beta, inv_beta);
y_vec[i] = o;
}
}
__global__ void softclip_kernel_scalar(
const float* __restrict__ x,
float* __restrict__ y,
int n,
float min_val,
float max_val,
float beta,
float inv_beta)
{
int idx = blockIdx.x * blockDim.x + threadIdx.x;
int stride = gridDim.x * blockDim.x;
for(int i = idx; i < n; i += stride) {
float v = x[i];
y[i] = v + softplus_f(min_val - v, beta, inv_beta) - softplus_f(v - max_val, beta, inv_beta);
}
}
torch::Tensor softclip_cuda(torch::Tensor x, float min_val, float max_val, float beta) {
auto x_c = x.contiguous();
auto output = torch::empty_like(x_c);
int n = x_c.numel();
int threads = 256;
float inv_beta = 1.0f / beta;
if (n % 4 == 0) {
int vec_n = n / 4;
int blocks = (vec_n + threads - 1) / threads;
if (blocks > 65535) blocks = 65535;
if (blocks == 0) blocks = 1;
softclip_kernel_ilp<<<blocks, threads>>>(
x_c.data_ptr<float>(),
output.data_ptr<float>(),
vec_n, min_val, max_val, beta, inv_beta
);
} else {
int blocks = (n + threads - 1) / threads;
if (blocks > 65535) blocks = 65535;
if (blocks == 0) blocks = 1;
softclip_kernel_scalar<<<blocks, threads>>>(
x_c.data_ptr<float>(),
output.data_ptr<float>(),
n, min_val, max_val, beta, inv_beta
);
}
return output;
}
"""
self.op = load_inline(
name="softclip_opt_v3_fix",
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=["softclip_cuda"],
extra_cuda_cflags=["-O3"],
verbose=False
)
def forward(self, x):
return self.op.softclip_cuda(x, self.min_val, self.max_val, self.beta)