changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
1.74 kB
# 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)