Files
yiliao2026/src/vlm_detect/vlm_backup/local_vlm_adapter_node.py

213 lines
7.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/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()