From 3ec6dfc3582430d25407372bd9e9ca5f17026a48 Mon Sep 17 00:00:00 2001 From: Kotaro Yoshimoto Date: Fri, 31 Jul 2026 09:15:36 +0900 Subject: [PATCH] fix(metrics): make clearance_p5_m reproducible tdigest.TDigest.compress() re-inserts centroids in pyudorandom order from an unseeded global random, so the same values give a different p5 every time -- four distinct values out of five builds in one process, relative spread 2.7e-03. Any A/B on this repository is affected. Seeding alone is not enough: merged_percentile also depends on input order, because t-digest merging is not associative. Both call sites already walk in sorted() order and a test pins that. getstate/seed/setstate is not re-entrant; both call sites are on the main thread and the invariant is documented. The reported percentile changes once. Past values were not comparable to each other either, so nothing is lost. 8 tests, all mutation-verified. --- scenario_generation/metrics/tdigest.py | 52 +++++-- .../tests/test_tdigest_determinism.py | 129 ++++++++++++++++++ 2 files changed, 171 insertions(+), 10 deletions(-) create mode 100644 scenario_generation/tests/test_tdigest_determinism.py diff --git a/scenario_generation/metrics/tdigest.py b/scenario_generation/metrics/tdigest.py index e4e6ea6d8..ca96fe2e4 100644 --- a/scenario_generation/metrics/tdigest.py +++ b/scenario_generation/metrics/tdigest.py @@ -6,6 +6,9 @@ from __future__ import annotations +import random +from contextlib import contextmanager + import numpy as np from tdigest import TDigest @@ -13,6 +16,29 @@ # a ``tdigests*.jsonl`` sidecar for multi-GPU clearance-p5 merge. TDIGEST_KEY = "_tdigest" +# Arbitrary, but must never change: the percentile a digest reports depends on it. +_RNG_SEED = 20260731 + + +@contextmanager +def _deterministic_rng(): + """Pin the global ``random`` state so that identical samples give an identical digest. + + ``tdigest`` 0.5.x draws from the global ``random`` module while building and compressing, + so without this the same input lands on different centroids and reports a different + percentile. The previous state is restored, so other users of ``random`` are unaffected. + + Not reentrant: concurrent digest builds would interleave the save/seed/restore and both + lose determinism. Both call sites are on the main thread; moving one into a worker needs a + lock here first. + """ + state = random.getstate() + random.seed(_RNG_SEED) + try: + yield + finally: + random.setstate(state) + def is_tdigest_key(key: str) -> bool: return key == TDIGEST_KEY or key.startswith(f"{TDIGEST_KEY}_") @@ -24,18 +50,24 @@ def tdigest_dict_from_values(values: np.ndarray) -> dict | None: finite = finite[np.isfinite(finite)] if finite.size == 0: return None - digest = TDigest() - digest.batch_update(finite.tolist()) - digest.compress() - return digest.to_dict() + with _deterministic_rng(): + digest = TDigest() + digest.batch_update(finite.tolist()) + digest.compress() + return digest.to_dict() def merged_percentile(digest_dicts: list[dict], percentile: float) -> float: - """Merge serialized digests and return an approximate percentile in ``[0, 100]``.""" + """Merge serialized digests and return an approximate percentile in ``[0, 100]``. + + Order-dependent: t-digest merging is not associative, so a caller pooling shards must feed + them in a stable order. + """ if not digest_dicts: return float("inf") - digest = TDigest() - for d in digest_dicts: - digest.update_from_dict(d) - digest.compress() - return float(digest.percentile(float(percentile))) + with _deterministic_rng(): + digest = TDigest() + for d in digest_dicts: + digest.update_from_dict(d) + digest.compress() + return float(digest.percentile(float(percentile))) diff --git a/scenario_generation/tests/test_tdigest_determinism.py b/scenario_generation/tests/test_tdigest_determinism.py new file mode 100644 index 000000000..ae4dd6e97 --- /dev/null +++ b/scenario_generation/tests/test_tdigest_determinism.py @@ -0,0 +1,129 @@ +"""The t-digest wrappers must be reproducible: ``clearance_p5_m`` is derived from them, and the +underlying ``tdigest`` package draws from the global ``random`` module. +""" + +from __future__ import annotations + +import json +import random + +import numpy as np + +from scenario_generation.metrics.tdigest import ( + merged_percentile, + tdigest_dict_from_values, +) + + +def _samples(n: int = 5108, seed: int = 0) -> np.ndarray: + """A clearance-like series: same size and shape as one closed-loop segment.""" + rng = np.random.default_rng(seed) + return np.abs(rng.normal(5.0, 3.0, n)) + + +def test_digest_from_identical_values_is_identical(): + values = _samples() + first = tdigest_dict_from_values(values) + for _ in range(7): + assert tdigest_dict_from_values(values) == first + + +def test_percentile_from_identical_values_is_identical(): + values = _samples() + got = {merged_percentile([tdigest_dict_from_values(values)], 5.0) for _ in range(8)} + assert len(got) == 1, f"clearance-p5 is not reproducible: {sorted(got)}" + + +def test_merge_of_several_digests_is_identical(): + """The DDP path merges one digest per shard; that merge must be reproducible too.""" + values = _samples() + shards = [tdigest_dict_from_values(values[i::4]) for i in range(4)] + got = {merged_percentile(shards, 5.0) for _ in range(8)} + assert len(got) == 1, f"merged clearance-p5 is not reproducible: {sorted(got)}" + + +def test_percentile_still_approximates_the_true_value(): + """Pinning the RNG must not have pinned it to a *wrong* answer.""" + values = _samples() + got = merged_percentile([tdigest_dict_from_values(values)], 5.0) + assert abs(got - float(np.percentile(values, 5.0))) < 0.01 + + +def test_global_random_state_is_not_disturbed(): + values = _samples(n=512) + + random.seed(123) + expected = [random.random() for _ in range(3)] + + random.seed(123) + tdigest_dict_from_values(values) + merged_percentile([tdigest_dict_from_values(values)], 5.0) + got = [random.random() for _ in range(3)] + + assert got == expected + + +def test_empty_input_still_returns_none_and_inf(): + assert tdigest_dict_from_values(np.array([])) is None + assert tdigest_dict_from_values(np.array([np.nan, np.inf])) is None + assert merged_percentile([], 5.0) == float("inf") + + +def _shard_files(out_dir, n_shards: int = 4, n_routes_per_shard: int = 3): + """Write per-rank segments_*/tdigests_* pairs the way a DDP run does.""" + out_dir.mkdir(parents=True, exist_ok=True) + rng = np.random.default_rng(7) + for rank in range(n_shards): + seg_lines, dig_lines = [], [] + for j in range(n_routes_per_shard): + route = f"route_{rank}_{j}" + vals = np.abs(rng.normal(5.0, 3.0, 400)) + digest = tdigest_dict_from_values(vals) + seg_lines.append( + json.dumps( + { + "route": route, + "object": { + "clearance_min_m": float(vals.min()), + "clearance_mean_m": float(vals.mean()), + "clearance_p5_m": float(np.percentile(vals, 5)), + "clearance_finite_steps": int(vals.size), + "miss_thresh_m": 0.5, + "collision_steps": 0, + "collision_count": 0, + "miss_steps": 0, + "miss_count": 0, + }, + } + ) + ) + dig_lines.append(json.dumps({"route": route, "object": digest})) + (out_dir / f"segments_{rank}.jsonl").write_text("\n".join(seg_lines) + "\n") + (out_dir / f"tdigests_{rank}.jsonl").write_text("\n".join(dig_lines) + "\n") + + +def test_shard_merge_reload_gives_the_same_p5(tmp_path): + from scenario_generation.closed_loop_eval import load_segment_rows_with_tdigests + from scenario_generation.metrics.tdigest import TDIGEST_KEY + + out_dir = tmp_path / "run" + _shard_files(out_dir) + + def p5_from_disk() -> float: + rows = load_segment_rows_with_tdigests(out_dir) + digests = [r["object"][TDIGEST_KEY] for r in rows] + return merged_percentile(digests, 5.0) + + got = {p5_from_disk() for _ in range(5)} + assert len(got) == 1, f"merged shard p5 is not reproducible: {sorted(got)}" + + +def test_shard_rows_load_in_a_deterministic_order(tmp_path): + from scenario_generation.closed_loop_eval import load_segment_rows_with_tdigests + + out_dir = tmp_path / "run" + _shard_files(out_dir) + order = [r["route"] for r in load_segment_rows_with_tdigests(out_dir)] + for _ in range(4): + assert [r["route"] for r in load_segment_rows_with_tdigests(out_dir)] == order + assert order == sorted(order), "shard files must be walked in sorted() order"