diff --git a/Dockerfile.ros2 b/Dockerfile.ros2 new file mode 100644 index 0000000..ee84b20 --- /dev/null +++ b/Dockerfile.ros2 @@ -0,0 +1,28 @@ +# ROS 2 Humble integration image for FlyGuard. +FROM ros:humble-ros-base-jammy + +ENV DEBIAN_FRONTEND=noninteractive \ + PYTHONDONTWRITEBYTECODE=1 \ + PYTHONUNBUFFERED=1 \ + PIP_NO_CACHE_DIR=1 + +WORKDIR /app + +# ROS message packages used directly by flyguard_ros2_node.py. +RUN apt-get update \ + && apt-get install -y --no-install-recommends \ + python3-pip \ + ros-humble-sensor-msgs-py \ + ros-humble-vision-msgs \ + ros-humble-visualization-msgs \ + && rm -rf /var/lib/apt/lists/* + +COPY requirements.txt ./ +RUN python3 -m pip install --no-cache-dir -r requirements.txt + +COPY flyguard ./flyguard +COPY artifacts ./artifacts + +# Parameters and ROS remappings are passed to `docker run` after `--ros-args`. +ENTRYPOINT ["/ros_entrypoint.sh"] +CMD ["python3", "flyguard/flyguard_ros2_node.py"] diff --git a/flyguard/flyguard_ros2_node.py b/flyguard/flyguard_ros2_node.py new file mode 100644 index 0000000..d3b9127 --- /dev/null +++ b/flyguard/flyguard_ros2_node.py @@ -0,0 +1,246 @@ +"""ROS 2 adapter for the FlyGuard lidar-obstacle detector. + +The processing core deliberately has no ROS dependency. This node adapts +``sensor_msgs/msg/PointCloud2`` to the core's ``flyguard.cdr.PointCloud2`` +contract and publishes final decisions for consumers and RViz. +""" + +from __future__ import annotations + +import math +import sys +from pathlib import Path + +import numpy as np + +ROOT_DIR = Path(__file__).resolve().parent.parent +if str(ROOT_DIR) not in sys.path: + sys.path.insert(0, str(ROOT_DIR)) + +import rclpy +from rclpy.executors import ExternalShutdownException +from rclpy.node import Node +from rclpy.qos import (QoSDurabilityPolicy, QoSHistoryPolicy, QoSProfile, + QoSReliabilityPolicy) + +from geometry_msgs.msg import Vector3 +from sensor_msgs.msg import PointCloud2 as RosPointCloud2 +import sensor_msgs_py.point_cloud2 as pc2 +from std_msgs.msg import String +from vision_msgs.msg import BoundingBox3D, Detection3D, Detection3DArray +from visualization_msgs.msg import Marker, MarkerArray + +from flyguard.cdr import PointCloud2 as CorePointCloud2 +from flyguard.export import export_frame +from flyguard.mbon_readout import MbonReadout +from flyguard.mushroom_body import MushroomBody +from flyguard.pipeline import FlyGuard, Params + + +# Matches the publisher profile recorded in the supplied rosbag: RELIABLE, +# VOLATILE, KEEP_LAST(10). PointCloud2 dimensions and payload size are read +# from every message, not configured in QoS. +LIDAR_QOS = QoSProfile( + history=QoSHistoryPolicy.KEEP_LAST, + depth=10, + reliability=QoSReliabilityPolicy.RELIABLE, + durability=QoSDurabilityPolicy.VOLATILE, +) + + +class FlyGuardNode(Node): + """Consume lidar clouds and publish final FlyGuard decisions.""" + + def __init__(self) -> None: + super().__init__("flyguard_node") + self.declare_parameter("lidar_topic", "/pandar_points") + # Empty means preserve the PointCloud2 header frame for downstream use. + self.declare_parameter("frame_id", "") + self.declare_parameter("memory_path", "artifacts/mushroom_body.npz") + self.declare_parameter("mbon_path", "artifacts/mbon_readout.npz") + self.declare_parameter("fov_deg", 30.0) + + self.lidar_topic = self.get_parameter("lidar_topic").value + self.frame_id_override = self.get_parameter("frame_id").value + memory_path = self.get_parameter("memory_path").value + mbon_path = self.get_parameter("mbon_path").value + fov_deg = self.get_parameter("fov_deg").value + + memory = self._load_model(memory_path, "memory", MushroomBody.load) + readout = self._load_model(mbon_path, "MBON readout", MbonReadout.load) + self.fg = FlyGuard(Params(fov_deg=float(fov_deg)), memory=memory, + readout=readout) + + # Match the supplied lidar publisher exactly; it is RELIABLE. + self.sub_cloud = self.create_subscription( + RosPointCloud2, self.lidar_topic, self.pointcloud_callback, + LIDAR_QOS, + ) + self.pub_threat = self.create_publisher(String, "flyguard/threat_level", 10) + self.pub_boxes = self.create_publisher( + Detection3DArray, "flyguard/bounding_boxes", 10) + self.pub_markers = self.create_publisher(MarkerArray, "flyguard/markers", 10) + + self.received_frames = 0 + self.processed_frames = 0 + self.get_logger().info( + "FlyGuard started: input=%s, frame=%s, memory=%s, mbon=%s" + % (self.lidar_topic, self.frame_id_override or "", + memory_path or "disabled", mbon_path or "disabled") + ) + + def _load_model(self, value: str, label: str, loader): + """Load an optional model and fail early for a supplied bad path.""" + if not value: + self.get_logger().warning(f"{label} is disabled (empty path)") + return None + path = Path(value) + if not path.is_absolute(): + path = ROOT_DIR / path + if not path.is_file(): + raise FileNotFoundError( + f"{label} file does not exist: {path}. Pass an empty parameter " + "only when the model is intentionally disabled." + ) + self.get_logger().info(f"Loading {label} from {path}") + return loader(path) + + @staticmethod + def _to_core_cloud(msg: RosPointCloud2) -> CorePointCloud2: + """Adapt ROS metadata and retain the complete raw scan ordering.""" + # Do not skip NaNs: removing no-return points changes the point order + # and prevents the core from recognising an organised Pandar scan. + points = np.asarray(pc2.read_points(msg, skip_nans=False)) + names = points.dtype.names or () + missing = {"x", "y", "z", "intensity"}.difference(names) + if missing: + raise ValueError("PointCloud2 is missing required fields: %s" % + ", ".join(sorted(missing))) + expected = int(msg.height) * int(msg.width) + if points.size != expected: + raise ValueError( + f"PointCloud2 has {points.size} decoded points, expected {expected}") + stamp = float(msg.header.stamp.sec) + float(msg.header.stamp.nanosec) * 1e-9 + return CorePointCloud2( + stamp=stamp, + frame_id=msg.header.frame_id, + height=int(msg.height), + width=int(msg.width), + point_step=int(msg.point_step), + is_dense=bool(msg.is_dense), + points=points, + ) + + def pointcloud_callback(self, msg: RosPointCloud2) -> None: + self.received_frames += 1 + started = self.get_clock().now() + try: + cloud = self._to_core_cloud(msg) + if cloud.n_points == 0: + self.get_logger().warning("Ignoring an empty PointCloud2 frame") + return + result = self.fg.process(cloud) + except Exception as exc: + # A malformed frame must not kill the ROS executor. + self.get_logger().error(f"FlyGuard failed to process cloud: {exc}") + return + + # The first calib_frames scans intentionally produce no result. + if result is None: + return + + self.processed_frames += 1 + decision = result.decision + frame_id = self.frame_id_override or msg.header.frame_id + exported = export_frame(decision, result.plane, result.corridor, + stamp=result.stamp) + threat_level = ("EMERGENCY" if decision.emergency else + "WARNING" if decision.detected else "CLEAR") + + self.pub_threat.publish(String(data=threat_level)) + self.publish_detections(exported.boxes, msg.header.stamp, frame_id) + self.publish_rviz_markers(exported.boxes, msg.header.stamp, frame_id) + + elapsed_ms = (self.get_clock().now() - started).nanoseconds / 1e6 + self.get_logger().debug( + "frame=%d processed=%d %.1f ms: %s, %d object(s)" + % (self.received_frames, self.processed_frames, elapsed_ms, + threat_level, len(exported.boxes)) + ) + + def publish_detections(self, boxes, stamp, frame_id: str) -> None: + message = Detection3DArray() + message.header.stamp = stamp + message.header.frame_id = frame_id + for box in boxes: + detection = Detection3D() + detection.header = message.header + detection.bbox = BoundingBox3D() + detection.bbox.center.position.x = box.x + detection.bbox.center.position.y = box.y + detection.bbox.center.position.z = box.z + detection.bbox.center.orientation.z = math.sin(box.yaw * 0.5) + detection.bbox.center.orientation.w = math.cos(box.yaw * 0.5) + detection.bbox.size = Vector3(x=box.dx, y=box.dy, z=box.dz) + message.detections.append(detection) + self.pub_boxes.publish(message) + + def publish_rviz_markers(self, boxes, stamp, frame_id: str) -> None: + message = MarkerArray() + clear = Marker() + clear.header.stamp = stamp + clear.header.frame_id = frame_id + clear.action = Marker.DELETEALL + message.markers.append(clear) + + for box in boxes: + cube = Marker() + cube.header.stamp = stamp + cube.header.frame_id = frame_id + cube.ns = "flyguard_boxes" + cube.id = int(box.track_id) + cube.type = Marker.CUBE + cube.action = Marker.ADD + cube.pose.position.x, cube.pose.position.y, cube.pose.position.z = box.x, box.y, box.z + cube.pose.orientation.z = math.sin(box.yaw * 0.5) + cube.pose.orientation.w = math.cos(box.yaw * 0.5) + cube.scale = Vector3(x=max(box.dx, 0.01), y=max(box.dy, 0.01), z=max(box.dz, 0.01)) + if box.threat_level.name == "EMERGENCY": + cube.color.r, cube.color.g, cube.color.b, cube.color.a = 1.0, 0.0, 0.0, 0.75 + else: + cube.color.r, cube.color.g, cube.color.b, cube.color.a = 1.0, 0.85, 0.0, 0.65 + message.markers.append(cube) + + label = Marker() + label.header.stamp = stamp + label.header.frame_id = frame_id + label.ns = "flyguard_labels" + label.id = 100000 + int(box.track_id) + label.type = Marker.TEXT_VIEW_FACING + label.action = Marker.ADD + label.pose.position.x, label.pose.position.y = box.x, box.y + label.pose.position.z = box.z + box.dz * 0.5 + 0.35 + label.pose.orientation.w = 1.0 + label.scale.z = 0.40 + label.color.r = label.color.g = label.color.b = label.color.a = 1.0 + label.text = "ID:%d | %.1f m | conf:%.2f" % ( + box.track_id, box.distance_along_track, box.confidence) + message.markers.append(label) + self.pub_markers.publish(message) + + +def main(args=None) -> None: + rclpy.init(args=args) + node = FlyGuardNode() + try: + rclpy.spin(node) + except (KeyboardInterrupt, ExternalShutdownException): + pass + finally: + node.destroy_node() + if rclpy.ok(): + rclpy.shutdown() + + +if __name__ == "__main__": + main()