forked from ccf-ai-infra/GPUCodeForces
198 lines
5.2 KiB
Python
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()
|
|
) |