From 10f39d10551f24d59aeb78fe98682c2a7ba079a4 Mon Sep 17 00:00:00 2001 From: Fabio Bagni <228594@studenti.unimore.it> Date: Tue, 4 May 2021 11:21:05 +0200 Subject: [PATCH] Fix tracker for batch size > 1 --- include/tkDNN/CenternetDetection3DTrack.h | 12 +- src/CenternetDetection3DTrack.cpp | 168 +++++++++++----------- 2 files changed, 92 insertions(+), 88 deletions(-) diff --git a/include/tkDNN/CenternetDetection3DTrack.h b/include/tkDNN/CenternetDetection3DTrack.h index 451bf6e..dd092fa 100644 --- a/include/tkDNN/CenternetDetection3DTrack.h +++ b/include/tkDNN/CenternetDetection3DTrack.h @@ -1,11 +1,11 @@ #ifndef CENTERNETDETECTION3DTRACK_H #define CENTERNETDETECTION3DTRACK_H +#include +#include "opencv2/opencv.hpp" #include "kernels.h" #include "utils.h" #include "tkdnn.h" -#include -#include "opencv2/opencv.hpp" #include #include #include // std::iota @@ -51,7 +51,7 @@ struct trackingRes class CenternetDetection3DTrack : public DetectionNN3D { -private: +public: tk::dnn::dataDim_t dim; tk::dnn::dataDim_t dim2; tk::dnn::dataDim_t dim_hm; @@ -148,9 +148,9 @@ private: std::vector det_res; int count_det; //tracks - std::vector tr_res; + std::vector> tr_res; std::vector> batchTracked; - int count_tr; + std::vector count_tr; int track_id=0; @@ -161,7 +161,7 @@ private: void pre_inf(const int bi); void _get_additional_inputs(); cv::Mat transform_preds_with_trans(float x1, float x2); - void tracking(); + void tracking(int bi); public: tk::dnn::Network *pre_phase_net = nullptr; diff --git a/src/CenternetDetection3DTrack.cpp b/src/CenternetDetection3DTrack.cpp index 02db674..25c4c0b 100644 --- a/src/CenternetDetection3DTrack.cpp +++ b/src/CenternetDetection3DTrack.cpp @@ -13,12 +13,13 @@ bool CenternetDetection3DTrack::init(const std::string& tensor_path, const int n nBatches = n_batches; confThreshold = conf_thresh; inputCalibs = k_calibs; + tr_res.resize(nBatches); init_preprocessing(); init_pre_inf(); init_postprocessing(); init_visualization(n_classes); - count_tr = 0; + count_tr.resize(nBatches, 0); } bool CenternetDetection3DTrack::init_preprocessing(){ @@ -30,6 +31,7 @@ bool CenternetDetection3DTrack::init_preprocessing(){ trans2 = cv::Mat(cv::Size(3,2), CV_32F); trans_out = cv::Mat(cv::Size(3,2), CV_32F); + dst2.at(0,0)=width * 0.5; dst2.at(0,1)=width * 0.5; dst2.at(1,0)=width * 0.5; @@ -372,6 +374,8 @@ void CenternetDetection3DTrack::preprocess(cv::Mat &frame, const int bi, const s sz = imageF.size(); cv::warpAffine(imageF, imageF, trans, cv::Size(dim.w, dim.h), cv::INTER_LINEAR ); + + cv::imshow("warp", imageF); sz = imageF.size(); imageF.convertTo(imageF, CV_32FC3, 1/255.0); @@ -414,7 +418,7 @@ cv::Mat CenternetDetection3DTrack::transform_preds_with_trans(float x1, float x2 return trans_out * target_coords; } -void CenternetDetection3DTrack::tracking(){ +void CenternetDetection3DTrack::tracking(int bi){ float item_size[count_det]; int item_cl[count_det]; @@ -427,44 +431,44 @@ void CenternetDetection3DTrack::tracking(){ dets[i*2+1] = det_res[i].ct.at(0,1); } - float track_size[count_tr]; - int track_cl[count_tr]; - float tracks[2*count_tr]; - for(int i=0; i(0,0) - tr_res[i].det_res.bb0.at(0,0)) * - (tr_res[i].det_res.bb1.at(0,1) - tr_res[i].det_res.bb0.at(0,1)); - track_cl[i] = tr_res[i].det_res.cl; - tracks[i*2] = tr_res[i].det_res.ct.at(0,0); - tracks[i*2+1] = tr_res[i].det_res.ct.at(0,1); + float track_size[count_tr[bi]]; + int track_cl[count_tr[bi]]; + float tracks[2*count_tr[bi]]; + for(int i=0; i(0,0) - tr_res[bi][i].det_res.bb0.at(0,0)) * + (tr_res[bi][i].det_res.bb1.at(0,1) - tr_res[bi][i].det_res.bb0.at(0,1)); + track_cl[i] = tr_res[bi][i].det_res.cl; + tracks[i*2] = tr_res[bi][i].det_res.ct.at(0,0); + tracks[i*2+1] = tr_res[bi][i].det_res.ct.at(0,1); } - float dist[count_tr*count_det]; + float dist[count_tr[bi]*count_det]; bool invalid; - for(int i=0; i track_size[i] || dist[j*count_tr+i] > item_size[j] || item_cl[j] != track_cl[i]; - dist[j*count_tr+i] = dist[j*count_tr+i] + invalid * (1 << 18); + invalid = dist[j*count_tr[bi]+i] > track_size[i] || dist[j*count_tr[bi]+i] > item_size[j] || item_cl[j] != track_cl[i]; + dist[j*count_tr[bi]+i] = dist[j*count_tr[bi]+i] + invalid * (1 << 18); } } - int matched_indices[2*count_tr]; + int matched_indices[2*count_tr[bi]]; float min_tr; int min_idtr=-1; - for(int i=0; i new_tr_res; int id_new_tr=0; - for(int i=0; i new_thresh) { count_tr_ ++; @@ -587,10 +591,10 @@ void CenternetDetection3DTrack::tracking(){ new_tr_res_.age = 1; new_tr_res_.active = 1; new_tr_res_.color = rand() % 256; - tr_res.push_back(new_tr_res_); + tr_res[bi].push_back(new_tr_res_); } } - count_tr = count_tr_; + count_tr[bi] = count_tr_; if(track_id==1000) track_id=0; @@ -729,8 +733,8 @@ void CenternetDetection3DTrack::postprocess(const int bi, const bool mAP) { } // track step - tracking(); - batchTracked.push_back(tr_res); + tracking(bi); + batchTracked.push_back(tr_res[bi]); } void CenternetDetection3DTrack::draw(std::vector& frames) {