213 lines
7.5 KiB
Python
213 lines
7.5 KiB
Python
#!/usr/bin/env python3
|
||
# -*- coding: utf-8 -*-
|
||
"""
|
||
Local VLM adapter for hobot_llamacpp.
|
||
|
||
This node keeps the vlm_detect external contract while delegating inference to
|
||
hobot_llamacpp over ROS topics.
|
||
"""
|
||
|
||
import threading
|
||
import time
|
||
|
||
import cv2
|
||
import numpy as np
|
||
import rclpy
|
||
from ai_msgs.msg import PerceptionTargets
|
||
from cv_bridge import CvBridge
|
||
from origincar_msg.srv import Speak
|
||
from rclpy.node import Node
|
||
from sensor_msgs.msg import CompressedImage, Image
|
||
from std_msgs.msg import Int32, String
|
||
|
||
from .local_adapter_utils import extract_perception_text, is_trigger_match
|
||
|
||
|
||
class LocalVLMAdapter(Node):
|
||
def __init__(self):
|
||
super().__init__("local_vlm_adapter")
|
||
|
||
self.declare_parameter("image_topic", "/image_mjpeg")
|
||
self.declare_parameter("trigger_topic", "/sign4return")
|
||
self.declare_parameter("trigger_sign", 9)
|
||
self.declare_parameter(
|
||
"prompt_text", "请描述这张图片的内容,用一句简短的话概括,不超过20个字。"
|
||
)
|
||
self.declare_parameter("result_topic", "/vlm_result")
|
||
self.declare_parameter("llamacpp_prompt_topic", "/prompt_text")
|
||
self.declare_parameter("llamacpp_image_topic", "/llamacpp/image")
|
||
self.declare_parameter("llamacpp_result_topic", "/llama_cpp_node")
|
||
self.declare_parameter("tts_service", "/tts/speak")
|
||
self.declare_parameter("enable_tts", True)
|
||
self.declare_parameter("inference_timeout_sec", 90.0)
|
||
self.declare_parameter("publish_delay_sec", 0.05)
|
||
|
||
image_topic = self.get_parameter("image_topic").value
|
||
trigger_topic = self.get_parameter("trigger_topic").value
|
||
self.trigger_sign = self.get_parameter("trigger_sign").value
|
||
self.prompt_text = self.get_parameter("prompt_text").value
|
||
result_topic = self.get_parameter("result_topic").value
|
||
self.inference_timeout_sec = float(
|
||
self.get_parameter("inference_timeout_sec").value
|
||
)
|
||
self.publish_delay_sec = float(self.get_parameter("publish_delay_sec").value)
|
||
self.enable_tts = bool(self.get_parameter("enable_tts").value)
|
||
|
||
llamacpp_prompt_topic = self.get_parameter("llamacpp_prompt_topic").value
|
||
llamacpp_image_topic = self.get_parameter("llamacpp_image_topic").value
|
||
llamacpp_result_topic = self.get_parameter("llamacpp_result_topic").value
|
||
tts_service = self.get_parameter("tts_service").value
|
||
|
||
self.bridge = CvBridge()
|
||
self.latest_image_msg = None
|
||
self.image_lock = threading.Lock()
|
||
self.inference_lock = threading.Lock()
|
||
self.waiting_for_result = False
|
||
self.pending_result = ""
|
||
self.result_event = threading.Event()
|
||
|
||
self.image_sub = self.create_subscription(
|
||
CompressedImage, image_topic, self.image_callback, 10
|
||
)
|
||
self.trigger_sub = self.create_subscription(
|
||
Int32, trigger_topic, self.trigger_callback, 10
|
||
)
|
||
self.llamacpp_result_sub = self.create_subscription(
|
||
PerceptionTargets, llamacpp_result_topic, self.llamacpp_result_callback, 10
|
||
)
|
||
|
||
self.prompt_pub = self.create_publisher(String, llamacpp_prompt_topic, 10)
|
||
self.image_pub = self.create_publisher(Image, llamacpp_image_topic, 10)
|
||
self.result_pub = self.create_publisher(String, result_topic, 10)
|
||
|
||
self.tts_client = self.create_client(Speak, tts_service)
|
||
|
||
self.get_logger().info(
|
||
"Local VLM adapter ready | image=%s | trigger=%s(sign=%s) | "
|
||
"prompt=%s | image_out=%s | result_in=%s | result_out=%s"
|
||
% (
|
||
image_topic,
|
||
trigger_topic,
|
||
self.trigger_sign,
|
||
llamacpp_prompt_topic,
|
||
llamacpp_image_topic,
|
||
llamacpp_result_topic,
|
||
result_topic,
|
||
)
|
||
)
|
||
|
||
def image_callback(self, msg):
|
||
with self.image_lock:
|
||
self.latest_image_msg = msg
|
||
|
||
def trigger_callback(self, msg):
|
||
if not is_trigger_match(msg.data, self.trigger_sign):
|
||
return
|
||
|
||
with self.inference_lock:
|
||
if self.waiting_for_result:
|
||
self.get_logger().warning("Local VLM inference already running")
|
||
return
|
||
self.waiting_for_result = True
|
||
self.pending_result = ""
|
||
self.result_event.clear()
|
||
|
||
thread = threading.Thread(target=self._run_inference_once, daemon=True)
|
||
thread.start()
|
||
|
||
def llamacpp_result_callback(self, msg):
|
||
text = extract_perception_text(msg)
|
||
if not text:
|
||
return
|
||
with self.inference_lock:
|
||
if not self.waiting_for_result:
|
||
return
|
||
self.pending_result = text
|
||
self.result_event.set()
|
||
|
||
def _run_inference_once(self):
|
||
try:
|
||
image_msg = self._build_image_msg_from_latest()
|
||
if image_msg is None:
|
||
self.get_logger().warning("No cached image available for local VLM")
|
||
return
|
||
|
||
prompt_msg = String()
|
||
prompt_msg.data = self.prompt_text
|
||
self.prompt_pub.publish(prompt_msg)
|
||
time.sleep(self.publish_delay_sec)
|
||
self.image_pub.publish(image_msg)
|
||
|
||
self.get_logger().info("Local VLM request sent to hobot_llamacpp")
|
||
if not self.result_event.wait(timeout=self.inference_timeout_sec):
|
||
self.get_logger().error(
|
||
"Timed out waiting for hobot_llamacpp result after %.1fs"
|
||
% self.inference_timeout_sec
|
||
)
|
||
return
|
||
|
||
with self.inference_lock:
|
||
result = self.pending_result
|
||
|
||
result_msg = String()
|
||
result_msg.data = result
|
||
self.result_pub.publish(result_msg)
|
||
self._speak_async(result)
|
||
self.get_logger().info("Local VLM result: %s" % result)
|
||
finally:
|
||
with self.inference_lock:
|
||
self.waiting_for_result = False
|
||
self.pending_result = ""
|
||
self.result_event.clear()
|
||
|
||
def _build_image_msg_from_latest(self):
|
||
with self.image_lock:
|
||
compressed = self.latest_image_msg
|
||
if compressed is None:
|
||
return None
|
||
|
||
np_arr = np.frombuffer(compressed.data, np.uint8)
|
||
cv_image = cv2.imdecode(np_arr, cv2.IMREAD_COLOR)
|
||
if cv_image is None:
|
||
self.get_logger().error("Failed to decode cached compressed image")
|
||
return None
|
||
|
||
image_msg = self.bridge.cv2_to_imgmsg(cv_image, encoding="bgr8")
|
||
image_msg.header = compressed.header
|
||
return image_msg
|
||
|
||
def _speak_async(self, text):
|
||
if not self.enable_tts:
|
||
return
|
||
if not self.tts_client.service_is_ready():
|
||
self.get_logger().warning("TTS service not available")
|
||
return
|
||
req = Speak.Request()
|
||
req.text = text
|
||
future = self.tts_client.call_async(req)
|
||
future.add_done_callback(self._tts_done_callback)
|
||
|
||
def _tts_done_callback(self, future):
|
||
try:
|
||
resp = future.result()
|
||
if not resp.success:
|
||
self.get_logger().warning("TTS failed: %s" % resp.message)
|
||
except Exception as exc:
|
||
self.get_logger().error("TTS call error: %s" % exc)
|
||
|
||
|
||
def main(args=None):
|
||
rclpy.init(args=args)
|
||
node = LocalVLMAdapter()
|
||
try:
|
||
rclpy.spin(node)
|
||
except KeyboardInterrupt:
|
||
pass
|
||
finally:
|
||
node.destroy_node()
|
||
rclpy.shutdown()
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|