changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
|
Raw History Blame Contribute Delete
13.5 kB

Python API: Diffusion Planner v5.0 (Autoware diffusion_planner) on Blackhole

Use this API from Python code (a pipeline, a notebook, a ROS 2 node wrapper). You do not need the HTTP server: the API and the server share the decoders, the device trace and the post-processing, so the outputs and the speed are the same.

Install

Install the package on top of an environment that already has ttnn (a tt-metal python_env at 44d66500520 with patches/tt-metal-eth-dispatch.patch, or the tt-model container). From the root of the model repository (the directory that holds pyproject.toml, README.md and code/):

pip install -e .                        # the Python API (numpy<2, pillow, pyyaml, onnx, huggingface_hub)
pip install -e ".[server,test]"         # + the HTTP server and the tests

The pip project is the repository's top-level pyproject.toml; it installs the package from code/tt_diffusion_planner (there is no pyproject.toml inside code/, because the container build copies code/ over the tt-metal tree). ttnn and torch come from tt-metal and are not declared.

The package carries tt_diffusion_planner.ttaw, the shared code of the Autoware ports to Blackhole (device open, trace runner, decoders, model base class, HTTP app), vendored at the version recorded in code/tt_diffusion_planner/ttaw/VENDORED.json.

You want to run Extras
the Python API none
the HTTP server (tt_diffusion_planner.server.app, see SERVING.md) server
host tests (no device; device tests are skipped): TT_VISIBLE_DEVICES=none python -m pytest -q code/tt_diffusion_planner/tests server,test
device tests: python -m pytest -q -s code/tt_diffusion_planner/tests/test_pcc_device.py code/tt_diffusion_planner/tests/test_e2e_device.py test

Quickstart

from tt_diffusion_planner import DiffusionPlanner

with DiffusionPlanner.from_pretrained(device_id=0) as model:
    out = model(inputs="code/tt_diffusion_planner/samples/kashiwanoha_dense.npz")
print(out.to_dict())   # the POST /predict body

examples/quickstart.py runs the same snippet, writes quickstart.json and a bird's-eye view of the input and the plan (quickstart_bev.png).

DiffusionPlanner.from_pretrained(...)

DiffusionPlanner.from_pretrained(
    model_id=None,          # HF repo or a local directory with the weights files; default AutowareFoundation/diffusion_planner
    *,
    revision=None,          # default for the default repo: the validated commit 423efde67f5 (tag v5.0)
    variant=None,           # "default" (the only v5.0 graph); default $DIFFUSION_PLANNER_VARIANT or "default"
    device_id=None,         # chip to open; default $TT_DEVICE_ID or 0
    device=None,            # an already-opened ttnn device (tt_diffusion_planner.device.open_device); close() does not close it
    dispatch=None,          # "eth" (p150 target, 12x10 grid) | "worker" (A/B only, 11x10) | "auto"; default $DIFFUSION_PLANNER_DISPATCH or "eth"
    num_command_queues=None,  # default $DIFFUSION_PLANNER_NUM_CQS or 1
    weights_dir=None,       # explicit local weights directory; no Hub access
    warmup_variants="default",  # trace variants to capture now; see "Warm-up"
    verbose=False,
    precision=None,         # the only compile parameter: extra precision-policy rules, e.g. "dec.*=HiFi2+fp32" (experiments only)
) -> DiffusionPlanner

What it does: resolves the weights first, so a Hub problem never claims the chip (weights_dir > $DIFFUSION_PLANNER_WEIGHTS_DIR > a local model_id directory > the HF snapshot at the pinned revision, restricted to the three v5.0 ONNX files and diffusion_planner.param.json, with an offline fallback to the cache; the sha256 of every file and the weights' major_version == 5 are checked), opens the chip (ETH dispatch, 12×10, 1 CQ; the other open parameters are DEVICE_DEFAULTS in tt_diffusion_planner/device.py, overridable with DIFFUSION_PLANNER_*), reads the ONNX initializers as data, uploads the weights and constants (48.1 MB), builds the graph, then compiles and captures the metal traces. If ETH dispatch cannot open (tt-metal without the patch), it warns and falls back to WORKER dispatch (model.info["device"]["fallback"] names it). Any other keyword argument is a TypeError.

The numerics are not arguments: the published configuration is the default of the DIFFUSION_PLANNER_LN_FP32, _SPLIT_MATMUL, _ATTN_MATMUL, _HIDDEN_FP32 and _ATTN_FP32_ACC knobs (tt_diffusion_planner.tt.config.KNOBS, pinned in tt-model.yaml serve.env). Setting one of them in the environment changes the graph and invalidates the accuracy figures of the card until the gates are re-run.

Warm-up

from_pretrained returns a warm model: it builds the graph, runs the plan once eagerly (this first run compiles every kernel into the JIT cache), then captures the whole plan as metal traces (warmup_variants="default": the full-capacity variant plan and one variant per agent bucket, plan_r32, plan_r64, plan_r96, plan_r128, plan_r192; each call replays the smallest that holds the scene, exact) with program-cache misses forbidden, so no later call compiles anything. model.warmup() is idempotent; warmup_variants="none" defers the capture to model.warmup().

Measured on the shipped sample (model.info["warmup_ms"], 2026-10-11): the load takes 352 s with an empty JIT cache and 12.9 s with a warm one (build 0.76 s: the ONNX initializers read and 118.9 MB of weights and constants uploaded; warm-up and capture of the 6 traces 7.7 s; the rest is the device open). The first call then takes 37 ms and the second 32 ms (an .npz path; the stage bench's steady state: 27 ms p50 for decoded arrays, 32 ms for an .npz path). The 6 traces hold 60.3 MB of DRAM (trace_region_size 192 MiB).

Call: model(...)

Argument Type Description
inputs mapping / .npz path / bytes / JSON envelope the 15 raw planner tensors (see "Input types")
velocity_smoothing_window int, 1..79, default 8 forward moving average of the trajectory velocity, in points
stopping_threshold float >= 0, default 0.3 force stop below this smoothed speed (m/s), when the ego moves
turn_indicator_keep_offset float, default -1.25 added to the KEEP logit before the turn-indicator decision
return_denoising_steps bool, default False add the ego row of the 11 solver iterates (out.meta["denoising_steps"], [11, 81, 4], the node's ~/debug/denoising_steps)

Input types

  • inputs=: the 15 raw tensors of the Autoware node's DiffusionPlannerCore::create_input_data() (batch 1, float32, ego base_link frame, BEFORE normalization; names and shapes in tt_diffusion_planner.INPUT_SCHEMA): a {name: array} mapping (numpy or torch), an .npz path or its bytes, or the /predict envelope {"format": "npz", "data": <base64>} / {"format": "json", "arrays": {...}}. Names, shapes and finite values are checked (InputError). tt_diffusion_planner.load_inputs(source) is the same decoder.
  • Any other input (points, images, calibration, ...) is refused (InputError).

What the caller keeps (the API is stateless)

One call is one independent plan. The Autoware node keeps state between plans; to reproduce it over a sequence of plans, the caller keeps the same state (SERVING.md 3.5 has the details):

  • The tensors. The node's pre-processing from ROS messages and the Lanelet2 map (per-UUID agent buffers and their 0.1 s resampling, the ego history, lane / route / polygon / line-string selection and encoding, traffic lights, speed limits, goal, turn-indicator report history) is not part of the bundle.

  • The turn-indicator hold window (turn_indicator_hold_duration, 1.0 s in the node's YAML). Each call decides with a fresh manager; apply the node's hold across calls with the node's own manager:

    from tt_diffusion_planner.host.postprocess import TurnIndicatorManager
    
    manager = TurnIndicatorManager()                     # hold 1.0 s, KEEP offset -1.25 (the node's YAML)
    out = model(inputs=tensors)
    decision = manager.evaluate(out.turn_indicator["logits"], stamp_s=now_s, prev_report=int(tensors["turn_indicators"][0, 30]))
    command = decision.command                           # the held command while less than 1.0 s has passed
    
  • The initial solver state sampled_trajectories (x_T, normalised space): zeros is the node's default (temperature: [0.0]); for a temperature > 0 send N(0, 1) x temperature; for the RTC prefix (delay_step > 0) put the previous plan into the ego row, slots t = 0 .. delay_step (x as (x - 10) / 20, y as y / 20, cos / sin as they are, in the current ego frame). delay is accepted and ignored (the node's multi-step mode never reads it).

  • The map frame. Outputs are in base_link; the node transforms them with the current ego pose to map.

Output

model(...) returns a tt_diffusion_planner.Output (= ttaw.outputs.Trajectory); out.to_dict() is exactly the POST /predict body (SERVING.md section 3.2).

field type meaning
poses float32 [80, 7] the ego trajectory at 0.1-8.0 s in base_link: x, y, yaw, cos, sin, velocity, acceleration (out.columns), post-processed like the node's ~/output/trajectory
turn_indicator dict command (0 NO_COMMAND, 1 DISABLE, 2 ENABLE_LEFT, 3 ENABLE_RIGHT), command_name, keep_selected, held (always false: no hold window), the 5 raw logits (NONE, DISABLE, LEFT, RIGHT, KEEP), the decision's probabilities
predicted_agents float32 [N, 80, 5] x, y, yaw, cos, sin of each non-empty neighbour row, in input order
meta dict predicted_agent_rows (the rows of predicted_agents), predicted_agent_columns, force_stop, time_from_start_s, valid_counts (the entities the encoder saw), and with return_denoising_steps the encoded denoising_steps
timing_ms dict preprocess, device (host tensors + H2D + replay + D2H), postprocess, total

out.to_dicts() gives one {x, y, yaw, cos, sin, velocity, acceleration} dict per trajectory point; out.to_dict("npz") adds the poses as a lossless base64 NPZ array.

Lifetime and information

  • model.close() releases the traces and the persistent device tensors and closes the chip if the model opened it; idempotent. with calls it for you; an unclosed model is closed when Python exits.
  • model.info: weights (repo, tag, revision, path), device (dispatch, grid, CQs, fallback), variant, warm variants, warm-up times, runtime parameter defaults, the input schema, the numerics options and precision policy in effect, the trace (variants, persistent inputs, trace buffers in MB).
  • Calls from several threads are safe: the device calls are serialised. One model per process per chip.

Speed

Warm calls, batch 1, ETH dispatch, 1 CQ, 12×10, the pinned numerics (code/scripts/bench.py, 100 iterations, 2026-10-11, the optimized release, on a shared host; p50, with p99 in brackets):

stage shipped sample kashiwanoha_dense (88 neighbours: bucket r96)
.npz decode + schema check (path inputs only) 5.76 (7.33) ms
host pre-processing (the node's normalization, the encoder's host features, masks, x_T) 2.45 (3.16) ms
packing the persistent trace inputs · ttnn host tensors · H2D 0.52 · 1.57 · 0.76 ms
device trace, one blocking plan 20.05 (20.08) ms
D2H (one packed read) 0.22 ms
host post-processing (trajectory, predicted paths, turn decision) 1.39 (1.84) ms
model(inputs=arrays) end to end 26.73 (27.48) ms
model(inputs=<.npz path>) 32.27 (33.61) ms
back-to-back replays (device time per plan) 20.00 ms = 50.0 plans/s

The device time depends on the agent bucket the scene needs (exact compaction: the smallest of 32 / 64 / 96 / 128 / 192 decoder rows that holds the ego and every valid neighbour row, else the full capacity). Back to back: r32 (≤ 31 neighbours, e.g. straight_road) 17.44 ms, r64 (a nuScenes instant with 42 neighbours) 19.25 ms, r96 20.00 ms, r128 21.11 ms, r192 22.30 ms, full capacity (> 191) 26.65 ms (the last three: OPT_REPORT.md round 5). The map entities are computed at full capacity in every plan. The first release took 102.04 ms for every scene (OPT_BASELINE.md). 2 CQs were not re-measured on this release (at the first release they did not help a synchronous request: OPT_BASELINE.md). Where the time goes and what comes next: OPT_REPORT.md.

Limits

  • Batch 1 on the chip; one model per process.
  • Fixed shapes of the v5.0 export (320 neighbours, 140 lanes, 25 route lanes, 10 polygons, 60 line strings, 31 history and 80 future steps) and 10 DPM-Solver steps (11 decoder evaluations) are compiled into the traces; the neighbour trunk and the decoder run on the smallest agent bucket that holds the scene (exact), so the device time depends on the neighbour count (17.4-26.7 ms); the map entities are computed at full capacity.
  • The node's guidance services (start / stop / centerline guidance) are not available (the node's default is off).
  • Accuracy is agreement with the fp32 CPU reference of the same network (README "Demo & Performances"); the planner's driving quality is the weights' (trained by TIER IV on data that is not public).