Update postprocess.cu

This commit is contained in:
Wang Xinyu 2023-11-14 18:36:17 +08:00 committed by GitHub
parent f39f2c1433
commit 2df7dd7f91
No known key found for this signature in database
GPG Key ID: 4AEE18F83AFDEB23

View File

@ -8,20 +8,21 @@ static __global__ void
decode_kernel(float *predict, int num_bboxes, float confidence_threshold, float *parray, int max_objects) {
float count = predict[0];
int position = (blockDim.x * blockIdx.x + threadIdx.x);
if (position >= count)
return;
if (position >= count) return;
float *pitem = predict + 1 + position * (sizeof(Detection) / sizeof(float));
int index = atomicAdd(parray, 1);
if (index >= max_objects)
return;
if (index >= max_objects) return;
float confidence = pitem[4];
if (confidence < confidence_threshold)
return;
if (confidence < confidence_threshold) return;
float left = pitem[0];
float top = pitem[1];
float right = pitem[2];
float bottom = pitem[3];
float label = pitem[5];
float *pout_item = parray + 1 + index * bbox_element;
*pout_item++ = left;
*pout_item++ = top;
@ -38,10 +39,8 @@ box_iou(float aleft, float atop, float aright, float abottom, float bleft, float
float ctop = max(atop, btop);
float cright = min(aright, bright);
float cbottom = min(abottom, bbottom);
float c_area = max(cright - cleft, 0.0f) * max(cbottom - ctop, 0.0f);
if (c_area == 0.0f)
return 0.0f;
if (c_area == 0.0f) return 0.0f;
float a_area = max(0.0f, aright - aleft) * max(0.0f, abottom - atop);
float b_area = max(0.0f, bright - bleft) * max(0.0f, bbottom - btop);
@ -49,28 +48,20 @@ box_iou(float aleft, float atop, float aright, float abottom, float bleft, float
}
static __global__ void nms_kernel(float *bboxes, int max_objects, float threshold) {
int position = (blockDim.x * blockIdx.x + threadIdx.x);
int count = bboxes[0];
// float count = 0.0f;
if (position >= count)
return;
if (position >= count) return;
float *pcurrent = bboxes + 1 + position * bbox_element;
for (int i = 1; i < count; ++i) {
for (int i = 0; i < count; ++i) {
float *pitem = bboxes + 1 + i * bbox_element;
if (i == position || pcurrent[5] != pitem[5]) continue;
if (pitem[4] >= pcurrent[4]) {
if (pitem[4] == pcurrent[4] && i < position)
continue;
if (pitem[4] == pcurrent[4] && i < position) continue;
float iou = box_iou(
pcurrent[0], pcurrent[1], pcurrent[2], pcurrent[3],
pitem[0], pitem[1], pitem[2], pitem[3]
pcurrent[0], pcurrent[1], pcurrent[2], pcurrent[3],
pitem[0], pitem[1], pitem[2], pitem[3]
);
if (iou > threshold) {
pcurrent[6] = 0;
return;
@ -83,12 +74,11 @@ void cuda_decode(float *predict, int num_bboxes, float confidence_threshold, flo
cudaStream_t stream) {
int block = 256;
int grid = ceil(num_bboxes / (float)block);
decode_kernel << <
grid, block, 0, stream >> > ((float *)predict, num_bboxes, confidence_threshold, parray, max_objects);
decode_kernel<<<grid, block, 0, stream>>>((float*)predict, num_bboxes, confidence_threshold, parray, max_objects);
}
void cuda_nms(float *parray, float nms_threshold, int max_objects, cudaStream_t stream) {
int block = max_objects < 256 ? max_objects : 256;
int grid = ceil(max_objects / (float)block);
nms_kernel << < grid, block, 0, stream >> > (parray, max_objects, nms_threshold);
nms_kernel<<<grid, block, 0, stream>>>(parray, max_objects, nms_threshold);
}