Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 23 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -97,3 +97,26 @@ uv run --package diffusion-planner python scripts/export/export_onnx.py \
The exporter validates the generated ONNX models with ONNX Runtime. For ROS 2 node
parameters, topics, model compatibility, and launch instructions, see
`ros2_ws/src/deps/autoware_universe/planning/autoware_ml_planner/README.md`.

## Shard datasets

`scripts/dataset/convert_h5_to_shards.py` repacks the frames of a Parquet index into a versioned
tar-shard dataset (one member per frame, one partition per rosbag) using the `planner-shards`
workspace package. Frames are staged in chunks through a scratch directory, packed, key-set,
scrubbed and spot-checked bit-exact against the H5 source:

```bash
uv run --package diffusion-planner python scripts/dataset/convert_h5_to_shards.py \
/data/diffusion_planner_h5/indexes/train.parquet /data/diffusion_planner_shards \
--tag train-v1 --staging-dir /dev/shm/shard-staging
```

Train from the shards with the `shards` dataloader config; the transforms are the same as the
H5 loader's and the loader is sharded per rank by the dataset itself:

```bash
uv run --package diffusion-planner python scripts/train/train.py \
train/dataloader@dataloader=shards \
dataloader.dataset.root=/data/diffusion_planner_shards \
dataloader.dataset.keyset_path=/data/diffusion_planner_shards/keysets/train-v1.parquet
```
51 changes: 51 additions & 0 deletions configs/train/dataloader/shards.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,51 @@
# Shard-backed alternative to dataloader/default.yaml: the same transforms, frames read from a
# versioned tar-shard dataset produced by scripts/dataset/convert_h5_to_shards.py.
# Select with: train/dataloader@dataloader=shards dataloader.dataset.root=... dataloader.dataset.keyset_path=...
defaults:
- /train/dataloader/transforms@_global_.transform_configs.start_decision_augmentation: start_decision_augmentation
- /train/dataloader/transforms@_global_.transform_configs.fix_stop_point: fix_stop_point
- /train/dataloader/transforms@_global_.transform_configs.pose_augmentation: pose_augmentation
- /train/dataloader/transforms@_global_.transform_configs.ilqr_refinement: ilqr_refinement
- /train/dataloader/transforms@_global_.transform_configs.speed_augmentation: speed_augmentation
- /train/dataloader/transforms@_global_.transform_configs.goal: goal
- /train/dataloader/transforms@_global_.transform_configs.ego_shape_augmentation: ego_shape_augmentation
- /train/dataloader/transforms@_global_.transform_configs.turn_indicator_augmentation: turn_indicator_augmentation
- /train/dataloader/transforms@_global_.transform_configs.traffic_light: traffic_light
- /train/dataloader/transforms@_global_.transform_configs.normalization: normalization
- _self_

_target_: diffusion_planner.data.build_shard_dataloader
dataset:
_target_: diffusion_planner.data.ShardPlannerDataset
# Dataset root written by convert_h5_to_shards.py (contains versions/, manifests, shards).
root: ???
# Version tag, or "latest".
version: latest
# Key-set Parquet selecting the frames to train on (written next to the dataset by the converter).
keyset_path: ???
# The per-rank plan needs the loader's batch and worker settings.
batch_size: ${dataloader.batch_size}
num_workers: ${dataloader.num_workers}
seed: ${seed}
shuffle: true
chunk_size: 256
# Fraction of samples the plan may duplicate to give every worker whole batches.
max_pad_fraction: 0.01
transforms:
- ${transform_configs.start_decision_augmentation}
- ${transform_configs.fix_stop_point}
- ${transform_configs.pose_augmentation}
- ${transform_configs.ilqr_refinement}
- ${transform_configs.speed_augmentation}
- ${transform_configs.goal}
- ${transform_configs.ego_shape_augmentation}
- ${transform_configs.turn_indicator_augmentation}
- ${transform_configs.traffic_light}
- ${transform_configs.normalization}
# This batch size applies to each process/GPU.
batch_size: 256
# Sequential shard reads need far fewer workers than per-frame H5 reads.
num_workers: 8
prefetch_factor: 4
pin_memory: true
drop_last: true
4 changes: 4 additions & 0 deletions packages/diffusion_planner/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ dependencies = [
"onnxruntime-gpu<1.24",
"onnxscript>=0.7.1",
"plotly>=6.9.0",
"planner-shards",
"pyarrow>=25.0.0",
"timm>=1.0.28",
"torch>=2.13.0",
Expand All @@ -25,3 +26,6 @@ dependencies = [
[build-system]
requires = ["uv_build>=0.11.0,<0.12.0"]
build-backend = "uv_build"

[tool.uv.sources]
planner-shards = { workspace = true }
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
PlannerDataset,
build_dataloader,
)
from .shard_planner_dataset import ShardPlannerDataset, build_shard_dataloader
from .transforms import (
FillUnknownTrafficLightFutures,
PlannerDataNormalizer,
Expand Down Expand Up @@ -34,8 +35,10 @@
"PlannerTurnIndicatorAugmentation",
"PoseAugmentationCase",
"PlannerDataset",
"ShardPlannerDataset",
"Transform",
"apply_pose_augmentation",
"build_dataloader",
"build_shard_dataloader",
"fill_unknown_traffic_light_futures",
]
Loading