Spaces:
Running
Running
File size: 5,036 Bytes
f7ebf60 49a08f0 f7ebf60 755ec7d f7ebf60 755ec7d f7ebf60 755ec7d f7ebf60 755ec7d f7ebf60 755ec7d f7ebf60 755ec7d 49a08f0 755ec7d f7ebf60 49a08f0 f7ebf60 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 | """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)
|