GUI: Major rework to predict_subscriber,

This commit is contained in:
2026-02-24 14:49:12 +01:00
parent 0b18942d21
commit ce11e2bb55
+78 -10
View File
@@ -1,14 +1,21 @@
import json
from aare.common.logger_config import setup_logger
import numpy as np
import zmq
from PySide6.QtCore import QThread, Signal
logger = setup_logger("aareGUI")
from PySide6.QtCore import QThread, Signal
from PySide6.QtGui import QImage, QPixmap
# If you need Bayer conversion like your SampleCameraThread did:
import cv2
class PredictionSubscriber(QThread):
# emits parsed JSON payload (dict with keys: time, frame_id, shape, boxes)
prediction = Signal(dict)
# NEW: emit the image that came with the prediction stream
image = Signal(QPixmap)
def __init__(self, pred_zmq_url: str, topic: bytes | str = b"", parent=None):
super().__init__(parent)
self._ctx = zmq.Context()
@@ -22,19 +29,80 @@ class PredictionSubscriber(QThread):
self._sock.connect(pred_zmq_url)
self.running = True
def _try_parse_json(self, part: bytes) -> dict | None:
try:
decoded = json.loads(part.decode("utf-8"))
return decoded if isinstance(decoded, dict) else None
except Exception:
return None
def _decode_image(self, header: dict, data: bytes) -> QPixmap | None:
"""
Supports:
- header["type"] == "uint8"
- header["shape"] == [H, W] (Bayer) -> converted to RGB
- header["shape"] == [H, W, 3] (RGB) -> used directly
"""
if not header or header.get("type") != "uint8":
return None
shape = header.get("shape")
if not shape or not isinstance(shape, (list, tuple)):
return None
arr = np.frombuffer(data, dtype=np.uint8)
if len(shape) == 2:
h, w = int(shape[0]), int(shape[1])
if arr.size != h * w:
return None
bayer = arr.reshape((h, w))
rgb = cv2.cvtColor(bayer, cv2.COLOR_BAYER_GB2RGB)
elif len(shape) == 3 and int(shape[2]) == 3:
h, w, c = int(shape[0]), int(shape[1]), int(shape[2])
if arr.size != h * w * c:
return None
rgb = arr.reshape((h, w, 3))
else:
return None
# If you still need the horizontal flip you had before:
rgb = rgb[:, ::-1, :].copy()
qimage = QImage(rgb.data, rgb.shape[1], rgb.shape[0], QImage.Format.Format_RGB888)
return QPixmap.fromImage(qimage)
def run(self):
while self.running:
try:
parts = self._sock.recv_multipart()
if not parts:
continue
# publisher sends either raw JSON or [topic, json]
payload_bytes = parts[-1]
try:
payload = json.loads(payload_bytes.decode("utf-8"))
except Exception:
continue
self.prediction.emit(payload)
json_dicts: list[dict] = []
non_json_parts: list[bytes] = []
# Parse all parts; keep non-JSON as candidates for image bytes
for p in parts:
d = self._try_parse_json(p)
if d is not None:
json_dicts.append(d)
else:
non_json_parts.append(p)
header = next((d for d in json_dicts if "shape" in d and d.get("type") == "uint8"), None)
detections = next((d for d in json_dicts if "boxes" in d), None)
# Heuristic: image payload is usually the largest non-JSON part
image_bytes = max(non_json_parts, key=len) if non_json_parts else None
if header and image_bytes:
pix = self._decode_image(header, image_bytes)
if pix is not None:
self.image.emit(pix)
if detections:
self.prediction.emit(detections)
except Exception as e:
logger.error("PredictionSubscriber error:", e)
break