armnet-eval / frames.py
villekuosmanen's picture
Ship armnet demo Space (policy eval on real robots)
f7ebf60 verified
Raw
History Blame Contribute Delete
3.39 kB
"""Turn a streamed RerunPacket into a single displayable video frame.
The eval runtime emits `RerunPacket` protobufs (one per control tick) carrying
each camera's image plus scalar joint/box readings. We decode the packet
directly — no Rerun viewer needed — pull out the image entries, and tile the
cameras side by side into one RGB frame for the Gradio image widget.
"""
from __future__ import annotations
import io
import logging
from typing import Optional
import numpy as np
logger = logging.getLogger(__name__)
# Tile each camera to this height (px), preserving aspect ratio, then hstack.
_TILE_HEIGHT = 288
# Cameras we prefer to show left-to-right; anything else follows, sorted.
_CAMERA_ORDER = ("front", "top", "left", "right", "wrist")
def decode_frame(data: bytes) -> Optional[np.ndarray]:
"""Decode one packet's images into a single tiled RGB uint8 frame, or None."""
from armnet_runtime import _rerun_pb2 as pb
from armnet_runtime.rerun import decode_packet
try:
packet = decode_packet(data)
except Exception: # noqa: BLE001 - a bad packet must never kill the stream
logger.debug("failed to decode rerun packet", exc_info=True)
return None
images: list[tuple[str, np.ndarray]] = []
for entry in packet.entries:
if entry.WhichOneof("value") != "image":
continue
arr = _decode_image(entry.image, pb)
if arr is not None:
images.append((entry.entity_path, arr))
if not images:
return None
images.sort(key=lambda pair: _camera_rank(pair[0]))
return _tile([arr for _, arr in images])
def _decode_image(image, pb) -> Optional[np.ndarray]: # noqa: ANN001
raw = bytes(image.data)
if not raw:
return None
if image.encoding == pb.IMAGE_ENCODING_JPEG:
try:
from PIL import Image as PILImage
return np.asarray(PILImage.open(io.BytesIO(raw)).convert("RGB"))
except Exception: # noqa: BLE001
logger.debug("failed to decode JPEG image entry", exc_info=True)
return None
if image.encoding == pb.IMAGE_ENCODING_RAW_RGB:
channels = image.channels or 3
arr = np.frombuffer(raw, dtype=np.uint8)
try:
if channels > 1:
return arr.reshape(image.height, image.width, channels)[..., :3]
gray = arr.reshape(image.height, image.width)
return np.stack([gray, gray, gray], axis=-1)
except ValueError:
logger.debug("raw image entry shape mismatch", exc_info=True)
return None
return None
def _camera_rank(entity_path: str) -> tuple[int, str]:
lowered = entity_path.lower()
for i, name in enumerate(_CAMERA_ORDER):
if name in lowered:
return (i, lowered)
return (len(_CAMERA_ORDER), lowered)
def _resize_to_height(arr: np.ndarray, height: int) -> np.ndarray:
from PIL import Image as PILImage
h, w = arr.shape[0], arr.shape[1]
if h == height:
return arr
new_w = max(1, round(w * height / h))
resized = PILImage.fromarray(arr).resize((new_w, height), PILImage.BILINEAR)
return np.asarray(resized)
def _tile(arrays: list[np.ndarray]) -> np.ndarray:
if len(arrays) == 1:
return arrays[0]
resized = [_resize_to_height(a, _TILE_HEIGHT) for a in arrays]
return np.hstack(resized)