GPUCodeForces/S1/uucoco_#118/CenterNetLoss_cuda.py

198 lines
5.2 KiB
Python

import torch
import torch.nn as nn
from torch.utils.cpp_extension import load_inline
cuda_source = """
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <math.h>
__inline__ __device__ float warp_reduce(float val) {
for (int offset = 16; offset > 0; offset /= 2)
val += __shfl_down_sync(0xffffffff, val, offset);
return val;
}
__global__ void centernet_hm_loss_kernel(
const float* __restrict__ pred_hm,
const float* __restrict__ gt_hm,
float* __restrict__ global_buffer,
int n_elements)
{
int tid = blockIdx.x * blockDim.x + threadIdx.x;
int stride = gridDim.x * blockDim.x;
float local_pos = 0.0f;
float local_neg = 0.0f;
float local_num = 0.0f;
for (int idx = tid; idx < n_elements; idx += stride) {
float p = pred_hm[idx];
float t = gt_hm[idx];
if (p < 1e-6f) p = 1e-6f;
if (p > 0.999999f) p = 0.999999f;
if (t == 1.0f) {
float term = (1.0f - p);
local_pos += logf(p) * term * term;
local_num += 1.0f;
} else {
float w = (1.0f - t);
w = w * w * w * w;
local_neg += logf(1.0f - p) * p * p * w;
}
}
local_pos = warp_reduce(local_pos);
local_neg = warp_reduce(local_neg);
local_num = warp_reduce(local_num);
if ((threadIdx.x % 32) == 0) {
atomicAdd(&global_buffer[0], local_pos);
atomicAdd(&global_buffer[1], local_neg);
atomicAdd(&global_buffer[2], local_num);
}
}
__global__ void centernet_reg_loss_kernel(
const float* __restrict__ pred,
const float* __restrict__ gt,
const float* __restrict__ mask,
float* __restrict__ global_buffer,
int output_idx,
int n_elements,
int spatial)
{
int tid = blockIdx.x * blockDim.x + threadIdx.x;
int stride = gridDim.x * blockDim.x;
float local_loss = 0.0f;
int stride_pred = 2 * spatial;
for (int idx = tid; idx < n_elements; idx += stride) {
int mask_idx = (idx / stride_pred) * spatial + (idx % spatial);
if (mask[mask_idx] == 1.0f) {
local_loss += fabsf(pred[idx] - gt[idx]);
}
}
local_loss = warp_reduce(local_loss);
if ((threadIdx.x % 32) == 0) {
atomicAdd(&global_buffer[output_idx], local_loss);
}
}
__global__ void centernet_final_loss_kernel(
float* __restrict__ buffer,
float* __restrict__ out)
{
if (threadIdx.x == 0) {
float pos_sum = buffer[0];
float neg_sum = buffer[1];
float num_pos = buffer[2];
float wh_sum = buffer[3];
float reg_sum = buffer[4];
float hm_loss = 0.0f;
if (num_pos > 0.0f) {
hm_loss = -(pos_sum + neg_sum) / num_pos;
wh_sum /= num_pos;
reg_sum /= num_pos;
} else {
hm_loss = -neg_sum;
}
out[0] = hm_loss + 0.1f * wh_sum + 1.0f * reg_sum;
}
}
torch::Tensor launch_centernet_loss(
torch::Tensor pred_hm, torch::Tensor gt_hm,
torch::Tensor pred_wh, torch::Tensor gt_wh,
torch::Tensor pred_reg, torch::Tensor gt_reg,
torch::Tensor mask)
{
auto options = pred_hm.options();
auto out = torch::empty({1}, options);
auto buffer = torch::zeros({5}, options);
int n_hm = pred_hm.numel();
int threads = 256;
int blocks_hm = (n_hm + threads - 1) / threads;
if (blocks_hm > 256) blocks_hm = 256;
centernet_hm_loss_kernel<<<blocks_hm, threads>>>(
pred_hm.data_ptr<float>(),
gt_hm.data_ptr<float>(),
buffer.data_ptr<float>(),
n_hm
);
int n_reg = pred_wh.numel();
int batch_size = pred_wh.size(0);
int spatial = pred_wh.size(2) * pred_wh.size(3);
int blocks_reg = (n_reg + threads - 1) / threads;
if (blocks_reg > 256) blocks_reg = 256;
centernet_reg_loss_kernel<<<blocks_reg, threads>>>(
pred_wh.data_ptr<float>(),
gt_wh.data_ptr<float>(),
mask.data_ptr<float>(),
buffer.data_ptr<float>(),
3,
n_reg,
spatial
);
centernet_reg_loss_kernel<<<blocks_reg, threads>>>(
pred_reg.data_ptr<float>(),
gt_reg.data_ptr<float>(),
mask.data_ptr<float>(),
buffer.data_ptr<float>(),
4,
n_reg,
spatial
);
centernet_final_loss_kernel<<<1, 1>>>(
buffer.data_ptr<float>(),
out.data_ptr<float>()
);
return out;
}
"""
cpp_source = """
torch::Tensor launch_centernet_loss(
torch::Tensor pred_hm, torch::Tensor gt_hm,
torch::Tensor pred_wh, torch::Tensor gt_wh,
torch::Tensor pred_reg, torch::Tensor gt_reg,
torch::Tensor mask);
"""
centernet_loss_module = load_inline(
name='centernet_loss_opt',
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=['launch_centernet_loss'],
verbose=False
)
class ModelNew(nn.Module):
def __init__(self):
super(ModelNew, self).__init__()
self.op = centernet_loss_module
def forward(self, pred_hm, gt_hm, pred_wh, gt_wh, pred_reg, gt_reg, mask):
return self.op.launch_centernet_loss(
pred_hm.contiguous(), gt_hm.contiguous(),
pred_wh.contiguous(), gt_wh.contiguous(),
pred_reg.contiguous(), gt_reg.contiguous(),
mask.contiguous()
)