armnet-eval / frames.py
villekuosmanen's picture
deploy demo space
49a08f0 verified
Raw
History Blame Contribute Delete
5.04 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
_MAX_DISPLAY_WIDTH = 900
# Cameras we prefer to show left-to-right; anything else follows, sorted.
_CAMERA_ORDER = ("front", "top", "left", "right", "wrist")
# Packets arrive at the policy's control rate, so a failure that affects one
# affects them all. Report the first at warning level and the rest at debug,
# rather than either spamming the log or (worse) staying silent about why the
# video never appeared.
_decode_failure_reported = False
def decode_frame(data: bytes) -> Optional[np.ndarray]:
"""Decode one packet's images into a single tiled RGB uint8 frame, or None.
Never raises. Returning None covers everything the caller can do nothing
about — a corrupt packet, a packet carrying only scalars, a missing decode
dependency — because the alternative is an exception on the video thread,
which takes the live feed down for the rest of the run.
"""
try:
# Imported here rather than at module scope so a Space missing the
# protobuf decode degrades to "no video" instead of failing to start.
from armnet_runtime import _rerun_pb2 as pb
from armnet_runtime.rerun import decode_packet
packet = decode_packet(data)
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 prepare_display_frame(_tile([arr for _, arr in images]))
except Exception: # noqa: BLE001 - a bad packet must never kill the stream
global _decode_failure_reported
if not _decode_failure_reported:
_decode_failure_reported = True
logger.warning(
"cannot decode the cell's video packets; live video will not "
"render for this run",
exc_info=True,
)
else:
logger.debug("failed to decode rerun packet", exc_info=True)
return None
def prepare_display_frame(frame: np.ndarray) -> np.ndarray:
"""Downscale once before sharing a frame across every connected viewer."""
width = frame.shape[1]
if width <= _MAX_DISPLAY_WIDTH:
return frame
from PIL import Image as PILImage
height = round(frame.shape[0] * _MAX_DISPLAY_WIDTH / width)
return np.asarray(
PILImage.fromarray(frame).resize(
(_MAX_DISPLAY_WIDTH, height), PILImage.BILINEAR
)
)
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)