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)