Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
54 changes: 53 additions & 1 deletion tests/entrypoints/openai_api/test_openpi_connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,7 +68,7 @@ def test_pack_and_unpack_round_trip_numpy_values():
assert decoded["nested"][0]["done"] == np.bool_(True)


def test_unpack_accepts_msgpack_numpy_marker_dicts():
def test_unpack_accepts_vllm_native_marker_dicts():
action = np.asarray([[1.0, 2.0]], dtype=np.float32)
payload = {
b"actions": {
Expand All @@ -88,6 +88,58 @@ def test_unpack_accepts_msgpack_numpy_marker_dicts():
np.testing.assert_allclose(decoded[b"actions"], np.asarray([[1.0, 0.0]], dtype=np.float32))


def test_unpack_accepts_msgpack_numpy_marker_dicts():
action = np.asarray([[1.0, 2.0]], dtype=np.float32)
payload = {
b"actions": {
b"nd": True,
b"type": action.dtype.str,
b"kind": b"",
b"shape": action.shape,
b"data": action.tobytes(),
}
}

decoded = openpi_connection._unpack_numpy(payload)

np.testing.assert_allclose(decoded[b"actions"], action)
assert decoded[b"actions"].dtype == np.float32


def test_unpack_rejects_msgpack_numpy_structured_array_markers():
values = np.zeros(2, dtype=[("a", "<i4"), ("b", "<f4")])
payload = {
b"nd": True,
b"type": values.dtype.descr,
b"kind": b"V",
b"shape": values.shape,
b"data": values.tobytes(),
}

with pytest.raises(ValueError, match="Unsupported dtype"):
openpi_connection._unpack_numpy(payload)


def test_unpack_msgpack_numpy_packed_observation():
msgpack_numpy = pytest.importorskip("msgpack_numpy")
obs = {
"observation/exterior_image_0_left": np.zeros((180, 320, 3), dtype=np.uint8),
"observation/joint_position": np.arange(7, dtype=np.float32),
}

decoded = openpi_connection._unpack(msgpack_numpy.packb(obs))

np.testing.assert_array_equal(
decoded["observation/exterior_image_0_left"],
obs["observation/exterior_image_0_left"],
)
np.testing.assert_allclose(
decoded["observation/joint_position"],
obs["observation/joint_position"],
)
assert decoded["observation/joint_position"].dtype == np.float32


def test_unpack_accepts_openpi_client_ndarray_markers():
image = np.arange(6, dtype=np.uint8).reshape(2, 3)
payload = {
Expand Down
16 changes: 14 additions & 2 deletions vllm_omni/entrypoints/openpi/connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,13 @@
vLLM-native markers:
ndarray -> {nd: true, type, kind, shape, data}
scalar -> {nd: false, type, kind, data}

`kind` is required on the `nd` markers because it is what separates them from a
plain user mapping, but its value is dialect-specific: the vLLM-native markers
carry the dtype kind character while the `msgpack-numpy` package leaves it empty
for plain dtypes and sets "V" for structured ones. Both spellings are accepted.
Scalars packed by `msgpack-numpy` omit `kind` entirely and are therefore left as
mappings.
"""

from __future__ import annotations
Expand Down Expand Up @@ -83,9 +90,14 @@ def _decode_vllm_numpy_marker(obj: dict[Any, Any]) -> Any:
if nd is _MISSING or dtype is _MISSING or kind is _MISSING or data is _MISSING:
return _MISSING

dtype_obj = np.dtype(_decode_marker_text(dtype))
kind_text = _decode_marker_text(kind)
if dtype_obj.kind != kind_text:
# Structured markers carry a dtype descriptor list in `type`, so reject them
# before `np.dtype()` sees something it cannot parse.
if kind_text == "V":
raise ValueError("Unsupported dtype: structured arrays")

dtype_obj = np.dtype(_decode_marker_text(dtype))
if kind_text not in ("", dtype_obj.kind):
raise ValueError(f"NumPy dtype marker kind mismatch: {dtype_obj.kind!r} != {kind_text!r}")
if dtype_obj.kind in ("V", "O", "c"):
raise ValueError(f"Unsupported dtype: {dtype_obj}")
Expand Down
Loading