changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
1.23 kB
#!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
"""Build (and optionally send) a ``POST /predict`` request for diffusion-planner-p150. Standard library only.
# the planner tensors (.npz holding the 15 raw tensors of INPUT_SCHEMA), with an optional runtime param
python3 code/tt_diffusion_planner/server/client.py \
--inputs code/tt_diffusion_planner/samples/kashiwanoha_dense.npz --param stopping_threshold=0.3 --out req.json
curl -s localhost:20000/predict -H 'Content-Type: application/json' -d @req.json
# send it directly and print the response
python3 code/tt_diffusion_planner/server/client.py \
--inputs code/tt_diffusion_planner/samples/kashiwanoha_dense.npz --url http://127.0.0.1:20000
The implementation is the vendored ``ttaw/server/client.py`` (C08); this file runs it from the model repository
with any Python 3.9+, without numpy and without installing the package. In Python, use
``tt_diffusion_planner.ttaw.server.client`` (``build_request``, ``post``, ``wait_ready``, ...).
"""
import runpy
from pathlib import Path
if __name__ == "__main__":
runpy.run_path(str(Path(__file__).resolve().parents[1] / "ttaw" / "server" / "client.py"), run_name="__main__")