Download code/tt_diffusion_planner/io.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 1.74 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/io.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/io.py
-
curl -L -o io.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/io.py
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) | |