Replace bridge augmentation core with fast numpy implementation - #276
yukkysaito wants to merge 2 commits into
Conversation
Port the reference implementation from trajectory_data_augmentation_test into data_augmentation_bridge_core.py and rebuild StatePerturbation on top of it. The public interface (constructor, __call__, augment, centric_transform) is unchanged, so train.py and other callers need no modification. Speed (same synthetic sample, past 2 s / future 8 s): - old torch implementation: mean 912 ms/sample (max 2451 ms) - new implementation: mean 73 ms/sample (max 200 ms), ~12x faster, including the new decorrelation features Core improvements ported: - binary-search feasibility search (assumes monotone feasibility in bridge times, coarse-scan fallback) instead of the linear M x N scan - regula falsi merge-point solver with a positions-only length metric - closed-form quintic (minimum-jerk) lateral profile coefficients New history-decorrelation features (both default on): - randomize_bridge_times: sample random feasible M/N above the minimum so the same past maps to varied recoveries - past_bump: one random sin^3 (C2-continuous) lateral bump on the past conditioning side only, constraint-checked with resampling Behavior changes: - when no feasible candidate exists, halve the offset/heading and retry (max 8 times) instead of emitting a constraint-violating sample; if still infeasible the sample is left unaugmented (aug_flag=False) - initial bridge times start at 0.1 s and the minimal feasible pair is searched, instead of fixed M=1.0 s / N=1.5 s Numpy RNG seeds are derived from the torch RNG, so set_seed-based reproducibility is preserved (covered by a test).
|
Since this is an independent change, I would like to incorporate only this into the main code for training; if the results are positive, I plan to merge it into the |
danielsanchezaran
left a comment
There was a problem hiding this comment.
We evaluated this PR fairly extensively (T4DEV-57842) by training from scratch on a small dataset with an identical recipe across arms and scoring lateral-recovery behaviour in closed loop. The result supports merging: this augmentation recovers from lateral displacement clearly better than the quintic augmentation currently in use, and the margin replicated across two independent training runs per arm. Numbers, video comparisons and the full method are on the wiki page "On augmentation for the Diffusion Planner IL".
A few things found along the way.
1. test_call_end_to_end_recenters_ego fails on this branch
FAILED diffusion_planner/tests/test_data_augmentation_bridge.py::test_call_end_to_end_recenters_ego
E KeyError: 'goal_pose'
Reproduced on this branch unmodified (1 failed, 5 passed). The other five tests call aug.augment(...), which never touches goal_pose; this is the only test that goes through __call__, which reaches centric_transform and indexes inputs["goal_pose"] unconditionally.
The fixture in _make_inputs() doesn't provide that key. One line fixes it:
"static_objects": torch.zeros(batch_size, 5, 10),
+ "goal_pose": torch.zeros(batch_size, 4),
}Verified: 6 passed with that added.
Worth deciding which side is wrong, though — the test is arguably right that the fixture should mirror a real batch, but it's also reasonable for centric_transform to tolerate a missing goal_pose rather than KeyError. Either fix is fine; right now the suite is red.
2. Stopped-vehicle heading
compute_forward_yaw recomputes heading from consecutive positions over the whole trajectory. For a stationary or near-stationary vehicle that displacement is localisation noise, so the recomputed heading is arbitrary — up to 180° wrong — and arctan2(0, 0) on exactly duplicated points yields 0. Any sample where the ego is stopped within the future window is affected.
We tried a fix (nearest valid GT yaw for near-stationary points) and it measurably reduced the error, but models trained with it recovered worse and we could not isolate why, so we are not proposing it. Flagging the defect rather than the patch.
3. The hardcoded wheelbase is a dead end — don't spend time on it
wheel_base = 2.75 is wrong for the vehicle these datasets come from, and it feeds steering_angle = arctan(yaw_rate * wb / |speed|) written into ego_current_state[8].
That channel is never read: the encoder takes ego history from ego_agent_past, and the decoder consumes only [:4] (pose) and [4:5] (longitudinal velocity). We confirmed this empirically — two training runs differing only in the wheelbase used came out bit-identical across all rollouts.
So it cannot affect model behaviour today, and there's no downstream consumer either, since augmentation only writes training data while inference gets steering from the vehicle. It becomes real the moment anyone makes the model read channel 8. Mentioning it mainly as a warning: it's a trap for anyone trying to validate a fix to those channels end-to-end, because correct and broken produce identical training runs.
What we would like to add on top, after this merges
We have a follow-up branch ready and would open it as a separate PR once this one lands, so the two stay independently reviewable:
- Bridge parameters as CLI flags. Currently the offset range, lateral-acceleration limit, bridge-time randomisation and past bump are fixed at construction; only
augment_probis plumbed through. Defaults unchanged, so this is inert unless a flag is passed. We need this because a wider offset range than the current default measured meaningfully better, and we'd like that to be configurable rather than a patch. - Removing numpy call overhead in the feasibility search.
np.gradientwas ~25% of runtime andnp.isclose~16% (while comparing two scalars). Replacing both with specialised 1-D equivalents gives 485 → 340 ms per augmented batch of 32, with output asserted bit-identical to numpy across both of its spacing branches, including its uniform-spacing reduction which uses different arithmetic. - Optional worker-side augmentation (
--augment_in_workers, default off). The solve is inherently per-sample, so running it in the DataLoader workers overlaps GPU compute instead of stalling the train loop. Draw order changes, so this one is distribution-equivalent rather than bit-identical: acceptance 30.9% vs 30.2%, applied-offset p90 1.043 vs 1.061, two-sample KS p=0.415.
End to end that took a 200-epoch run from 9 h 22 min to 4 h 52 min. That matters here specifically because this augmentation kept improving out to ~200 epochs in our runs, where the quintic one plateaued much earlier — so the longer budget is the operating point, not an edge case.
Happy to reorder any of this if you'd rather fold some of it into this PR instead.
danielsanchezaran
left a comment
There was a problem hiding this comment.
Trained this from scratch against the current quintic augmentation and against no augmentation, same recipe everywhere, and scored 8-second lateral recovery in closed loop. It beats quintic clearly and the margin held across two runs per arm, so this is worth merging. Numbers and video comparisons are on the wiki page "On augmentation for the Diffusion Planner IL".
One red test and two bugs inline. Neither bug blocks the merge in my view — the wheelbase one is inert and we are deliberately not proposing a fix for the heading one. The test failure should be fixed before merging.
After this lands we would like to follow up with a separate PR: bridge params as CLI flags (defaults unchanged), and two opt-in throughput changes that took a 200-epoch run from 9h22 to 4h52. Relevant here because this augmentation kept improving out to ~200 epochs in our runs while quintic plateaued early, so the long budget is the normal case.
| "route_lanes": torch.zeros(batch_size, 5, 20, 12), | ||
| "polygons": torch.zeros(batch_size, 5, 8, 3), | ||
| "line_strings": torch.zeros(batch_size, 5, 8, 3), | ||
| "static_objects": torch.zeros(batch_size, 5, 10), |
There was a problem hiding this comment.
This fixture has no goal_pose, so test_call_end_to_end_recenters_ego fails:
E KeyError: 'goal_pose'
The other 5 tests call augment(), which never touches it. This is the only one going through __call__ -> centric_transform, which indexes inputs["goal_pose"] directly.
| "static_objects": torch.zeros(batch_size, 5, 10), | |
| "static_objects": torch.zeros(batch_size, 5, 10), | |
| "goal_pose": torch.zeros(batch_size, 4), |
With that, 6 passed. Alternatively make centric_transform skip a missing goal_pose instead — your call, but the suite is red right now.
| x: np.ndarray, y: np.ndarray, fallback_yaw: np.ndarray | None = None | ||
| ) -> np.ndarray: | ||
| sigma = cumulative_distance(x, y) | ||
| if len(x) >= 3 and sigma[-1] > 1.0e-6: |
There was a problem hiding this comment.
Heading gets recomputed from consecutive positions here. When the ego is stopped, that displacement is just localisation noise, so the yaw comes out arbitrary — up to 180 deg off — and duplicated points give arctan2(0, 0) = 0. The sigma[-1] > 1.0e-6 guard only catches a trajectory that is stopped everywhere, not one that stops partway through the future, which is the common case.
We tried passing nearest valid GT yaw for the near-stationary points. It cut the error a lot, but models trained on it recovered worse and we could not work out why, so we are not proposing that patch. Flagging the bug, not the fix.
| def __init__( | ||
| self, | ||
| augment_prob: float = 0.5, | ||
| wheel_base: float = 2.75, |
There was a problem hiding this comment.
This default is wrong for the vehicle our datasets come from, but before anyone fixes it: it only feeds steering_angle into ego_current_state[8], and nothing reads that channel. Encoder takes ego history from ego_agent_past, decoder reads [:4] and [4:5] only.
We checked by training twice, changing only the wheelbase — bit-identical across all rollouts.
So it is harmless today and there is no downstream consumer either, since augmentation only writes training data and inference gets steering from the vehicle. Worth knowing because it is a trap: correct and broken produce identical runs, so you cannot validate a fix here end to end. It only starts mattering if someone makes the model read channel 8.
| @@ -0,0 +1,968 @@ | |||
| """Numpy core for the bridge-based state perturbation augmentation. | |||
There was a problem hiding this comment.
Two hot spots in this file, if you want them: np.gradient was ~25% of augmentation runtime and np.isclose ~16% (comparing two scalars). Swapping both for 1-D versions gives 485 -> 340 ms per augmented batch of 32, bit-identical output.
We have that plus CLI flags for the bridge params ready on a branch — happy to send it as a separate PR once this lands, rather than piling it on here.
Intuitive visualizatin no the augmented trajectory after history bump ->
But notice that the heading offset is not continuous on the 0 point?
@SakodaShintaro cc: @yukkysaito Is this an expected update or a bug? |
| max_heading_offset_deg: float = 10.0, | ||
| max_lateral_accel_mps2: float = 3.0, | ||
| max_bridge_speed_gap_mps: float = 0.5, | ||
| max_bridge_jerk_mps3: float = 5.0, |
There was a problem hiding this comment.
| max_bridge_jerk_mps3: float = 5.0, | |
| max_bridge_jerk_mps3: float = 20.0, |
@SakodaShintaro @yukkysaito
From my local test using parts of our bus driving data. I found that almost all of the trajectory failed the jerk computation because native noise in the system.
This will leads to
- False negative on the evaluation of the old pipeline and fixed to M, N to default
- Longer search time.
- New pipeline collapses into M=N=0.1 (where the jerk/curvature computation is wrong and luckily pass the test)
After relaxing this, the augmented result started to look correct.
There was a problem hiding this comment.
In fact, I have barely touched the "bridge" mode for data augmentation, and it hasn't been used in training either. I'll take a closer look at it now.
There was a problem hiding this comment.
@SakodaShintaro Thanks for reminding me, so we have been training with the DP's default augmentation right?
There was a problem hiding this comment.
we have been training with the DP's default augmentation
Yes. I only use that mode.



Summary
Replaces the torch-based
data_augmentation_bridgecore with the fast numpy implementation developed intrajectory_data_augmentation_test, and adds two history-decorrelation features. TheStatePerturbationpublic interface (constructor,__call__,augment,centric_transform,_low/_high) is unchanged —train.pyand other callers need no modification.Performance (same synthetic sample, past 2 s / future 8 s)
Speedups ported from the reference repo:
History decorrelation (both default on)
The smooth past bridge encodes the offset and recovery timing in the ego history, letting the model shortcut map-based recovery by extrapolating its own history. Two mitigations:
randomize_bridge_times: sample random feasible M/N above the minimal feasible pair, so the same past maps to varied recoveriespast_bump: one randomsin^3(C2-continuous) lateral bump injected into the past conditioning side only, constraint-checked with resampling; sampling ranges are fixed constants indata_augmentation_bridge_core.pyBehavior changes
aug_flag=False) — the old implementation emitted the constraint-violating sample in this case.Numpy RNG seeds are derived from the torch RNG, so
set_seed-based reproducibility is preserved.Test plan
diffusion_planner/tests/test_data_augmentation_bridge.py(6 tests, all pass): shapes/continuity at t0, offset range, constraint feasibility of the output, slow-vehicle gating, determinism undertorch.manual_seed, end-to-end__call__withcentric_transform, runtime boundtest_data_augmentation.pyfailures (32) confirmed unrelated: identical failures ondevwithout this changeReference implementation and design docs: https://github.com/yukkysaito/trajectory_data_augmentation_test (PR #3)