File size: 1,739 Bytes
4d9b003 | 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 | # SPDX-License-Identifier: Apache-2.0
"""Input decoding and output encoding shared by the Python API and the HTTP server of diffusion-planner-p150.
The implementation is the vendored ``ttaw.io`` (C08); numpy only, data parsing only, and every client mistake raises
:class:`InputError`, which the server maps to HTTP 400. This model's only input is ``inputs``: the 15 ONNX-named
planner tensors of the node's ``create_input_data()`` (``reference.config.INPUT_SCHEMA``), as a ``{name: array}``
mapping, an ``.npz`` path or its bytes, or the JSON envelope ``{"format": "npz", "data": <base64>}`` /
``{"format": "json", "arrays": {...}}``. :func:`load_inputs` / :func:`decode_inputs` bind the schema, so
``model(inputs=...)`` and ``POST /predict`` accept and refuse the same inputs (names, shapes, finite values).
"""
from __future__ import annotations
from typing import Any, Mapping, Optional
from .reference.config import INPUT_SCHEMA
from .ttaw import io as _io
from .ttaw.io import * # noqa: F401,F403 (the decoders and encoders: ttaw/API.md section 9)
__all__ = list(_io.__all__) + ["DEFAULT_POINT_FIELDS", "INPUT_SCHEMA", "load_inputs", "decode_inputs"]
DEFAULT_POINT_FIELDS = () # no point cloud input
def load_inputs(source: Any) -> dict:
"""Python-API planner input (mapping, ``.npz`` path or bytes, JSON envelope) -> ``{name: float32 array}``
checked against ``INPUT_SCHEMA``."""
return _io.load_named_arrays(source, INPUT_SCHEMA)
def decode_inputs(spec: Mapping[str, Any], *, max_bytes: Optional[int] = None) -> dict:
"""The ``inputs`` envelope of ``/predict`` -> ``{name: float32 array}`` checked against ``INPUT_SCHEMA``."""
return _io.decode_named_arrays(spec, INPUT_SCHEMA, max_bytes=max_bytes)
|