diff --git a/src/target_observation_processing.cpp b/src/target_observation_processing.cpp new file mode 100644 index 0000000..a17c82f --- /dev/null +++ b/src/target_observation_processing.cpp @@ -0,0 +1,447 @@ +#include "target_observation_processing.hpp" + +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include + +namespace odin_ros_driver { + +namespace { + +constexpr std::array kHipKeypointIndices = {11, 12}; + +float median_in_place(std::vector& values) +{ + if (values.empty()) { + return -1.0f; + } + const auto mid = values.begin() + static_cast(values.size() / 2); + std::nth_element(values.begin(), mid, values.end()); + return *mid; +} + +std::string format_target_text(const TargetObservation& observation) +{ + std::ostringstream oss; + oss.setf(std::ios::fixed); + oss.precision(2); + oss << "ID:" << observation.track_id + << " conf:" << observation.confidence + << " depth:" << observation.depth << "m"; + return oss.str(); +} + +} // namespace + +void TargetObservationProcessor::initialize(const TargetObservationConfig& config) +{ + config_ = config; + if (config_.yolo_engine_path.empty()) { + throw std::runtime_error("Target observation YOLO engine path is empty"); + } + if (!config_.yolo_labels_path.empty() && + !std::filesystem::exists(config_.yolo_labels_path)) { + config_.yolo_labels_path.clear(); + } + + yolo_ = std::make_unique( + config_.yolo_engine_path, + config_.yolo_labels_path); + tracker_ = std::make_unique( + 0.3f, 30, 50, 3, 0.3f, false, 1); + current_target_id_ = -1; + initialized_ = true; +} + +TargetObservation TargetObservationProcessor::process( + const cv::Mat& camera_bgr, + const pcl::PointCloud& cloud_in_cam, + const CloudReprojector& reprojector, + const CloudReprojector::OdomPose& odom_pose, + TargetObservationDebugInfo* debug_info) +{ + const auto total_start = std::chrono::steady_clock::now(); + TargetObservation empty; + if (debug_info) { + *debug_info = TargetObservationDebugInfo{}; + debug_info->current_target_id_before = current_target_id_; + } + if (!initialized_ || camera_bgr.empty()) { + if (debug_info) { + debug_info->total_ms = + std::chrono::duration(std::chrono::steady_clock::now() - total_start).count(); + } + return empty; + } + + const auto yolo_start = std::chrono::steady_clock::now(); + const auto poses = yolo_->detect(camera_bgr, config_.yolo_conf, config_.yolo_nms); + const auto yolo_end = std::chrono::steady_clock::now(); + if (debug_info) { + debug_info->poses_count = static_cast(poses.size()); + debug_info->yolo_ms = + std::chrono::duration(yolo_end - yolo_start).count(); + } + if (poses.empty()) { + if (debug_info) { + debug_info->total_ms = + std::chrono::duration(std::chrono::steady_clock::now() - total_start).count(); + } + return empty; + } + + const Eigen::MatrixXf detections = format_detections(poses); + const auto mot_start = std::chrono::steady_clock::now(); + const Eigen::MatrixXf tracks = tracker_->update(detections, camera_bgr); + const auto mot_end = std::chrono::steady_clock::now(); + if (debug_info) { + debug_info->tracks_count = tracks.rows(); + debug_info->mot_ms = + std::chrono::duration(mot_end - mot_start).count(); + } + if (tracks.rows() == 0) { + if (debug_info) { + debug_info->total_ms = + std::chrono::duration(std::chrono::steady_clock::now() - total_start).count(); + } + return empty; + } + + TargetObservation target = select_target( + poses, tracks, camera_bgr.cols, camera_bgr.rows, debug_info); + if (!target.valid) { + if (debug_info) { + debug_info->total_ms = + std::chrono::duration(std::chrono::steady_clock::now() - total_start).count(); + } + return empty; + } + + target = enrich_target_with_cloud( + target, poses, cloud_in_cam, reprojector, odom_pose, debug_info); + if (!target.valid) { + if (debug_info) { + debug_info->total_ms = + std::chrono::duration(std::chrono::steady_clock::now() - total_start).count(); + } + return empty; + } + + current_target_id_ = target.track_id; + if (debug_info) { + debug_info->total_ms = + std::chrono::duration(std::chrono::steady_clock::now() - total_start).count(); + } + return target; +} + +Eigen::MatrixXf TargetObservationProcessor::format_detections( + const std::vector& poses) const +{ + Eigen::MatrixXf dets(poses.size(), 6); + for (size_t i = 0; i < poses.size(); ++i) { + dets(static_cast(i), 0) = poses[i].box.x; + dets(static_cast(i), 1) = poses[i].box.y; + dets(static_cast(i), 2) = poses[i].box.x + poses[i].box.width; + dets(static_cast(i), 3) = poses[i].box.y + poses[i].box.height; + dets(static_cast(i), 4) = poses[i].conf; + dets(static_cast(i), 5) = poses[i].classId; + } + return dets; +} + +TargetObservation TargetObservationProcessor::select_target( + const std::vector& poses, + const Eigen::MatrixXf& tracks, + int image_width, + int image_height, + TargetObservationDebugInfo* debug_info) +{ + TargetObservation observation; + int best_row = -1; + bool found_existing = false; + const float img_cx = static_cast(image_width) * 0.5f; + const float img_cy = static_cast(image_height) * 0.5f; + float best_dist = std::numeric_limits::max(); + + for (int i = 0; i < tracks.rows(); ++i) { + const int track_id = static_cast(tracks(i, 4)); + if (track_id == current_target_id_) { + best_row = i; + found_existing = true; + break; + } + + const float cx = 0.5f * (tracks(i, 0) + tracks(i, 2)); + const float cy = 0.5f * (tracks(i, 1) + tracks(i, 3)); + const float dist = (cx - img_cx) * (cx - img_cx) + (cy - img_cy) * (cy - img_cy); + if (dist < best_dist) { + best_dist = dist; + best_row = i; + } + } + + if (best_row < 0) { + return observation; + } + + const int detection_index = static_cast(tracks(best_row, 7)); + if (detection_index < 0 || + detection_index >= static_cast(poses.size())) { + return observation; + } + + const auto& pose = poses[static_cast(detection_index)]; + observation.valid = true; + if (debug_info) { + debug_info->found_existing_target = found_existing; + debug_info->target_selected = true; + } + observation.track_id = static_cast(tracks(best_row, 4)); + observation.detection_index = detection_index; + if (debug_info) { + debug_info->selected_track_id = observation.track_id; + debug_info->detection_index = detection_index; + } + observation.confidence = pose.conf; + observation.bbox_xyxy = { + static_cast(pose.box.x), + static_cast(pose.box.y), + static_cast(pose.box.x + pose.box.width), + static_cast(pose.box.y + pose.box.height), + }; + + for (size_t i = 0; i < std::min(pose.keypoints.size(), 17); ++i) { + observation.keypoints_xyc[i * 3 + 0] = pose.keypoints[i].x; + observation.keypoints_xyc[i * 3 + 1] = pose.keypoints[i].y; + observation.keypoints_xyc[i * 3 + 2] = pose.keypoints[i].confidence; + } + + if (!found_existing) { + current_target_id_ = observation.track_id; + } + return observation; +} + +TargetObservation TargetObservationProcessor::enrich_target_with_cloud( + const TargetObservation& target, + const std::vector& poses, + const pcl::PointCloud& cloud_in_cam, + const CloudReprojector& reprojector, + const CloudReprojector::OdomPose& odom_pose, + TargetObservationDebugInfo* debug_info) const +{ + const auto depth_start = std::chrono::steady_clock::now(); + if (!target.valid || + target.detection_index < 0 || + target.detection_index >= static_cast(poses.size()) || + cloud_in_cam.empty()) { + if (debug_info) { + debug_info->depth_ms = + std::chrono::duration(std::chrono::steady_clock::now() - depth_start).count(); + } + return TargetObservation{}; + } + + const auto& cam_params = reprojector.getCameraParams(); + const auto& pose = poses[static_cast(target.detection_index)]; + + std::vector uv_points; + std::vector cam_points; + uv_points.reserve(cloud_in_cam.size()); + cam_points.reserve(cloud_in_cam.size()); + + for (const auto& pt : cloud_in_cam) { + if (pt.z <= config_.min_depth || pt.z >= config_.max_depth) { + continue; + } + + const Eigen::Vector2d uv = reprojector.projectCameraPointToPixel( + Eigen::Vector3d(pt.x, pt.y, pt.z)); + if (uv.x() < 0.0 || uv.x() >= cam_params.image_width || + uv.y() < 0.0 || uv.y() >= cam_params.image_height) { + continue; + } + + uv_points.emplace_back( + static_cast(uv.x()), + static_cast(uv.y())); + cam_points.emplace_back(pt.x, pt.y, pt.z); + } + if (debug_info) { + debug_info->projected_cloud_points = static_cast(cam_points.size()); + } + + if (uv_points.empty()) { + if (debug_info) { + debug_info->depth_ms = + std::chrono::duration(std::chrono::steady_clock::now() - depth_start).count(); + } + return TargetObservation{}; + } + + std::vector per_keypoint_depths; + per_keypoint_depths.reserve(kHipKeypointIndices.size()); + + for (const int kp_index : kHipKeypointIndices) { + if (kp_index >= static_cast(pose.keypoints.size())) { + continue; + } + const auto& kp = pose.keypoints[static_cast(kp_index)]; + if (kp.confidence < 0.5f) { + continue; + } + + std::vector nearby_depths; + nearby_depths.reserve(32); + for (size_t i = 0; i < uv_points.size(); ++i) { + const float du = uv_points[i].x - kp.x; + const float dv = uv_points[i].y - kp.y; + if (std::hypot(du, dv) < config_.search_radius_px) { + nearby_depths.push_back(cam_points[i].z()); + } + } + if (nearby_depths.size() < 3) { + continue; + } + per_keypoint_depths.push_back(median_in_place(nearby_depths)); + } + if (debug_info) { + debug_info->depth_sample_count = static_cast(per_keypoint_depths.size()); + } + + if (per_keypoint_depths.empty()) { + if (debug_info) { + debug_info->depth_ms = + std::chrono::duration(std::chrono::steady_clock::now() - depth_start).count(); + } + return TargetObservation{}; + } + + if (per_keypoint_depths.size() >= 3) { + std::vector deviations = per_keypoint_depths; + const float median_depth = median_in_place(deviations); + for (size_t i = 0; i < deviations.size(); ++i) { + deviations[i] = std::abs(deviations[i] - median_depth); + } + const float mad = median_in_place(deviations); + if (mad > 1e-3f) { + const float threshold = 2.5f * mad / 0.6745f; + std::vector inliers; + inliers.reserve(per_keypoint_depths.size()); + for (const float depth : per_keypoint_depths) { + if (std::abs(depth - median_depth) < threshold) { + inliers.push_back(depth); + } + } + if (!inliers.empty()) { + per_keypoint_depths.swap(inliers); + } + } + } + + TargetObservation enriched = target; + std::vector final_depths = per_keypoint_depths; + enriched.depth = median_in_place(final_depths); + enriched.depth_confidence = std::min( + 1.0f, + static_cast(per_keypoint_depths.size()) / + static_cast(kHipKeypointIndices.size())); + + std::vector valid_hip_pixels; + valid_hip_pixels.reserve(kHipKeypointIndices.size()); + for (const int kp_index : kHipKeypointIndices) { + if (kp_index >= static_cast(pose.keypoints.size())) { + continue; + } + const auto& kp = pose.keypoints[static_cast(kp_index)]; + if (kp.confidence >= 0.5f) { + valid_hip_pixels.emplace_back(kp.x, kp.y); + } + } + if (valid_hip_pixels.empty()) { + const float center_u = 0.5f * (target.bbox_xyxy[0] + target.bbox_xyxy[2]); + const float center_v = 0.5f * (target.bbox_xyxy[1] + target.bbox_xyxy[3]); + enriched.target_pos_cam = reprojector.pixelToCameraPoint( + center_u, center_v, enriched.depth).cast(); + } else { + Eigen::Vector2f mean_pixel = Eigen::Vector2f::Zero(); + for (const auto& pixel : valid_hip_pixels) { + mean_pixel += pixel; + } + mean_pixel /= static_cast(valid_hip_pixels.size()); + enriched.target_pos_cam = reprojector.pixelToCameraPoint( + mean_pixel.x(), mean_pixel.y(), enriched.depth).cast(); + } + + if (enriched.target_pos_cam.z() <= 0.0f) { + if (debug_info) { + debug_info->depth_ms = + std::chrono::duration(std::chrono::steady_clock::now() - depth_start).count(); + } + return TargetObservation{}; + } + + enriched.target_pos_world = reprojector.cameraPointToWorld( + enriched.target_pos_cam.cast(), odom_pose).cast(); + enriched.valid = true; + if (debug_info) { + debug_info->depth_ms = + std::chrono::duration(std::chrono::steady_clock::now() - depth_start).count(); + } + return enriched; +} + +void draw_target_observation_overlay( + cv::Mat& image_bgr, + const TargetObservation& observation) +{ + if (image_bgr.empty() || !observation.valid) { + return; + } + + const cv::Scalar box_color(80, 220, 255); + const int x1 = static_cast(std::lround(observation.bbox_xyxy[0])); + const int y1 = static_cast(std::lround(observation.bbox_xyxy[1])); + const int x2 = static_cast(std::lround(observation.bbox_xyxy[2])); + const int y2 = static_cast(std::lround(observation.bbox_xyxy[3])); + cv::rectangle(image_bgr, cv::Point(x1, y1), cv::Point(x2, y2), box_color, 2); + + for (size_t i = 0; i < 17; ++i) { + const float conf = observation.keypoints_xyc[i * 3 + 2]; + if (conf <= 0.3f) { + continue; + } + const int u = static_cast(std::lround(observation.keypoints_xyc[i * 3 + 0])); + const int v = static_cast(std::lround(observation.keypoints_xyc[i * 3 + 1])); + cv::circle(image_bgr, cv::Point(u, v), 4, cv::Scalar(0, 255, 0), -1); + } + + const std::string label = format_target_text(observation); + cv::putText( + image_bgr, + label, + cv::Point(x1, std::max(20, y1 - 8)), + cv::FONT_HERSHEY_SIMPLEX, + 0.55, + cv::Scalar(0, 0, 0), + 3); + cv::putText( + image_bgr, + label, + cv::Point(x1, std::max(20, y1 - 8)), + cv::FONT_HERSHEY_SIMPLEX, + 0.55, + cv::Scalar(255, 255, 255), + 1); +} + +} // namespace odin_ros_driver