Lidar_Muxa/flyguard/flyguard_ros2_node.py
2026-09-26 21:47:02 +03:00

235 lines
9.7 KiB
Python

"""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 qos_profile_sensor_data
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
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)
# Pandar and most lidar drivers offer BEST_EFFORT sensor QoS. A
# default RELIABLE subscriber is incompatible and would receive no data.
self.sub_cloud = self.create_subscription(
RosPointCloud2, self.lidar_topic, self.pointcloud_callback,
qos_profile_sensor_data,
)
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 "<input header>",
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()