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)