Skip to content

Replace bridge augmentation core with fast numpy implementation - #276

Open
yukkysaito wants to merge 2 commits into
devfrom
feature/fast-bridge-augmentation
Open

yukkysaito wants to merge 2 commits into
devfrom
feature/fast-bridge-augmentation

Conversation

@yukkysaito

Copy link
Copy Markdown
Collaborator

Summary

Replaces the torch-based data_augmentation_bridge core with the fast numpy implementation developed in trajectory_data_augmentation_test, and adds two history-decorrelation features. The StatePerturbation public interface (constructor, __call__, augment, centric_transform, _low/_high) is unchanged — train.py and other callers need no modification.

Performance (same synthetic sample, past 2 s / future 8 s)

implementation mean max
old torch core 912 ms/sample 2451 ms
new numpy core (incl. new features) 73 ms/sample 200 ms

Speedups ported from the reference repo:

  • binary-search feasibility search over bridge times (assumes monotone feasibility, coarse-scan fallback for non-monotone pockets) instead of the linear M×N grid scan
  • regula falsi merge-point solver with a positions-only path-length metric
  • closed-form quintic (minimum-jerk) lateral profile coefficients

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 recoveries
  • past_bump: one random sin^3 (C2-continuous) lateral bump injected into the past conditioning side only, constraint-checked with resampling; sampling ranges are fixed constants in data_augmentation_bridge_core.py

Behavior changes

  • When no feasible candidate exists (short 2 s history can make large offsets infeasible), the offset/heading are halved and retried (max 8 times). If still infeasible the sample is left unaugmented (aug_flag=False) — the old implementation emitted the constraint-violating sample in this case.
  • 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.

Test plan

  • New 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 under torch.manual_seed, end-to-end __call__ with centric_transform, runtime bound
  • 50-sample benchmark: mean 73 ms, max 200 ms
  • Pre-existing test_data_augmentation.py failures (32) confirmed unrelated: identical failures on dev without this change

Reference implementation and design docs: https://github.com/yukkysaito/trajectory_data_augmentation_test (PR #3)

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).
@SakodaShintaro

Copy link
Copy Markdown

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 tier4-main branch.

@danielsanchezaran danielsanchezaran left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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_prob is 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.gradient was ~25% of runtime and np.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 danielsanchezaran left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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),

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Suggested change
"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:

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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,

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

@Owen-Liuyuxuan

Owen-Liuyuxuan commented Aug 10, 2026 •

Copy link
Copy Markdown
image

Intuitive visualizatin no the augmented trajectory after history bump ->

  • duration 1~2s
  • amplitude 0.05-0.15 including plus and minus.
  • increasing history decorrelation

But notice that the heading offset is not continuous on the 0 point?

image image

@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,

@Owen-Liuyuxuan Owen-Liuyuxuan Aug 10, 2026 •

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
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.

Image Image

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@SakodaShintaro Thanks for reminding me, so we have been training with the DP's default augmentation right?

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

we have been training with the DP's default augmentation

Yes. I only use that mode.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants