diff --git a/FRESH_RESTART.md b/FRESH_RESTART.md new file mode 100644 index 0000000..d6635e2 --- /dev/null +++ b/FRESH_RESTART.md @@ -0,0 +1,141 @@ +# Double-shift SFM fresh restart + +## Frozen scientific timepoint + +The quoted double-shift OOD statistics were produced from the clean source +commit `ca7f0d718f8d70cf74833b1c75157caf7f1b13f2` on July 20, 2026. The +authenticated environment-and-visualization child commit is +`b27df76fe461fd7f7e86ecf87e80cbf52e7f01d5`; this is the restart base. + +The benchmark contract is: + +- scene: `double_density_velocity_ood`; +- pedestrians: 40; +- pedestrian desired-speed range: 1.0--2.0 m/s; +- episode bank: 250000--250099; +- evaluation: raw temperature 1, NFE 8, M=100 per gamma; +- pretrained checkpoint SHA-256: + `1b5179c935d3eeff8824967d707d64cc9bab273949ee1f0e4f190172bab1b215`; +- historical A-r10 checkpoint SHA-256: + `bf6f521dd2dd6de4cffcce672a8ce4adbf00bb14e71dd9fd27704d205f65744c`. + +The authenticated pooled OOD results were: + +| Method | SR | CR | Successful clearance | Successful time | +|---|---:|---:|---:|---:| +| Hp10 r0 raw | 70.00% | 30.00% | 0.131 m | 8.69 s | +| historical A-r10 raw | 69.43% | 30.29% | 0.128 m | 8.79 s | +| Kazuki default | 81.29% | 18.71% | 0.181 m | 4.39 s | +| Kazuki goal-stress | 75.29% | 24.71% | 0.163 m | 3.81 s | + +The original metrics artifact is +`/home/dohyun/projects/sfm_hp10_b1_runs/ca7f0d7_preexp_double_shift/double_shift_ood/metrics.json` +on Helios. Its local Mac copy is +`/Users/dhl/Documents/SFM_HP10_DOUBLE_SHIFT_PREEXP_b27df76/double_shift_ood/metrics.json` +with SHA-256 +`566ac3fc87b727ad0957b837aca68a1fdd24777040584791bd372eb25e3b8977`. + +## Repository relationship + +`DHLeexpress/safe_flow_expansion_SFM` has a standalone packaging history, so +Git cannot express it as being a number of commits ahead of this historical +safeMPPI branch. Its current source snapshot records safeMPPI commit +`e5ab47b`, which is eight safeMPPI source commits after `b27df76` and changes +34 SFM files. Therefore this restart is published on isolated archive/restart +branches; `master` is not reset or force-pushed. + +## One deliberate correction after the timepoint + +The historical model, scene, exact K=16 moving-obstacle SOCP, and checkpoint are +preserved. One user-approved semantic correction is applied before new +experiments: every queried action window is certified over all H=10 +transitions. Crossing the goal does not truncate a queried window. Goal reach +only terminates the closed-loop episode after the selected first action. + +Thus new source should be described as **b27df76 plus the full-H10 correction**, +not as bitwise historical b27df76. + +## Pre-expansion full-episode diagnostic + +`sfm_b1_full_episode_audit.py` is an isolated diagnostic. It does not modify +the historical fail-closed trainer or enter any sample into D, D+, the GP, or +gradient replay. + +For three fixed episodes and all seven gamma values it starts with the +pretrained generator and the round-1 B1 mechanism: + +1. generate K=16 windows; +2. select B=4 using the empty-buffer RBF acquisition with pending-point + conditioning; +3. run the exact full-H10 verifier on B; +4. if an admissible B query exists, execute the max-one-step-Hp-margin action; +5. on finite-B NVP, independently sample and verify one raw temperature-1 + window, execute its first action, and continue the offline simulator; +6. stop only at realized collision, goal reach, or T=180 timeout. + +Step 5 is deliberately **not** a certified controller. It exists only to +observe post-NVP states that fail-closed gathering would hide. + +Labels remain separate: + +- `verifier_positive` / `verifier_negative`: exact safety label of the + actually executed H=10 window; +- `finite_B_NVP`: B=4 failed to contain an admissible candidate; it is not + itself a verifier-negative label; +- `trap`: displacement over the last ten executed transitions is below 0.2 m; +- `collision`: realized simulator outcome; it does not retroactively relabel + earlier safe windows. + +Consequently, NVP, trap, and collision must not be pooled blindly into the +negative verifier loss. They can support a later, separately defined +continuation/viability label. + +The companion `sfm_b1_full_episode_viz.py` renders: + +- K=16 generated paths in gray; +- B=4 queried paths in orange; +- verifier-positive/rejected B paths in green/red; +- the complete executed trail in blue/red according to the exact H=10 label; +- the executed positive candidate's exact K=16 verifier polytope and H=1..10 + level sets in green; +- finite-B NVP, first trap entry, and collision as separate markers. + +The resulting movie is a **pretrained-generator round-1 gathering diagnostic**, +not a pure raw-policy rollout and not a safety-certified deployment. + +## Authenticated round-1 diagnostic result + +The diagnostic was run from clean commit +`bf53dee110f885b28a7783db432085e7d75f15ff` on Helios GPU 3 using episodes +250001, 250003, and 250007 for every gamma. Collection took 2 min 57 s. + +- B queries: 5,198 verifier-positive and 2,410 verifier-negative; +- executed windows: 1,459 verifier-positive and 443 verifier-negative; +- finite-B NVP contexts: 454; +- NVP continuations: 11 certified raw rescues and 443 uncertified raw actions; +- episode outcomes: 18 success and 3 collision; +- first ten-step trap entries: 3. + +This is the central pre-expansion observation: a fail-closed controller would +have hidden 454 post-NVP states. Only 11 independently sampled raw windows at +those states were both full-H positive and nominal-Hp admissible. Continuing +the other 443 states exposes unsafe data, but does not make it certified data. + +Because this is round 1, the GP history buffer is empty. The RBF length scale +does not change the equal marginal prior variance; it affects only +pending-point conditioning among the K candidates. Therefore the reported +negative selected-versus-marginal uplift is not a round-2 novelty result and +must not be used to judge the historical RBF buffer. + +Server artifacts: + +`/data3/research1/sfm_fresh_b27df76_full_episode_audit_bf53dee` + +Mac artifacts: + +`/Users/dhl/Documents/SFM_HP10_FRESH_RESTART_B27DF76_BF53DEE` + +The video is H.264, 2436x1060, 75 frames, 15 s. Its SHA-256 is +`e2abc447ced326ebdfe6be155989ac2dae3f0fb856e6ba426c59fbd9163319d8`. +The full trace SHA-256 is +`eba53f8d389e6caf30569b31749ab4bb66e88d7e805e831623bf8f8589a392bb`. diff --git a/overnight_run_07_12_sfm/CLAUDE_NEUTRAL_EXPANSION_HANDOFF.md b/overnight_run_07_12_sfm/CLAUDE_NEUTRAL_EXPANSION_HANDOFF.md new file mode 100644 index 0000000..2324c40 --- /dev/null +++ b/overnight_run_07_12_sfm/CLAUDE_NEUTRAL_EXPANSION_HANDOFF.md @@ -0,0 +1,138 @@ +# Claude handoff: Hp10 neutral Safe Flow Expansion + +## Start here + +Claude must begin by running this command on Helios and showing its complete +stdout before editing or launching anything: + +```bash +cd /home/dohyun/projects/ +/home/dohyun/miniforge3/envs/cfm_mppi/bin/python \ + overnight_run_07_12_sfm/show_claude_neutral_handoff.py +``` + +The first line is the authenticated **Neutral continuation MP4**: + +```text +/data3/research1/sfm_neutral_claude_handoff_assets/neutral_continuation.mp4 +``` + +The command only reads and authenticates existing artifacts. It does not run an +evaluation. The cached reference delivery is +`/data3/research1/sfm_neutral_gamma_temp_0e441d6/DELIVERY_COMPLETE.json`. + +## Immutable shared inputs + +- Pretrained checkpoint: + `/home/dohyun/projects/sfm_hp10_b1_runs/103476d/pretrained_hp10.pt` +- Checkpoint SHA-256: + `1b5179c935d3eeff8824967d707d64cc9bab273949ee1f0e4f190172bab1b215` +- ID demonstration dataset: + `/home/dohyun/projects/cfm_mppi/overnight_run_07_12_sfm/dataset_id_v01` +- Pretraining report: + `/home/dohyun/projects/sfm_hp10_b1_runs/103476d/pretrain/pretraining_report.json` +- Existing neutral baseline run, for read-only comparison: + `/data3/research1/sfm_neutral_multiround_round50_a36dfe7/lr1em5_s01` +- New Claude outputs must be under a new root such as + `/data3/research1/claude_sfm_neutral_/`. Never reuse or modify a + shared result root. + +## ID SafeMPPI pretraining contract + +The demonstration distribution is 20 pedestrians with desired speeds sampled +in 0.5--1.0 m/s. Demonstrations are successful-only trajectories at all seven +gamma values `(0.1, 0.2, 0.3, 0.4, 0.5, 0.7, 1.0)`. + +`sfm_b1_expert.py` defines the faithful data expert: horizon 10, 2048 MPPI +rollouts, MPPI temperature 0.1, velocity-aware K=16 nominal polytope, +`centroid_gain=.2`, `centroid_smooth=.25`, `centroid_eps=.15`, +`predict_gain=.25`, and `smooth_weight=.12`. Temperature 0.1 belongs only to +the demonstration expert; it is not silently used by raw learned-policy +evaluation. + +`stage3_pretrain_sfm.py` constructs the ten-step Hp history, splits complete +trajectories 90/10, and trains from scratch with objective mass +`gamma -> successful trajectory -> window`. Its defaults are 120 epochs, +batch 256, AdamW lr `3e-4`, weight decay `1e-4`, five warmup epochs, cosine +decay, and seed `20260720`. Checkpoint promotion uses ID validation followed by +a fixed trajectory-disjoint ID raw temperature-1 gate; no OOD result selects +the pretrained checkpoint. + +## Canonical neutral expansion recipe + +The entry point is `sfm_b1_neutral_multiround.py`. The handoff baseline is the +explicit `lr1em5_s01` recipe, not the script's historical smoke defaults: + +- OOD collector: 40 pedestrians, 1.0--2.0 m/s. +- One macro-round: two new scenario seeds x all seven gamma values, synchronously. +- `T=180`, `H=10`, `K=16`, `B=4`, NFE 8, generation temperature 1. +- Penultimate noised representation at `s=.9`; each stored plan retains its + original Gaussian base `x0`. +- RBF-GP: previous round's executed full-H positives only, gamma-balanced cap + 512, `ell=.24210826720721101`, lambda `1e-2`, adaptive ESS target `.5`. + The GP is fixed within the macro-round and rebuilt once at the next round. +- Exact moving-pedestrian full-H verifier labels every queried/repaired plan. +- If ordinary queried candidates are admissible, execute the full-H positive + passing the nominal-Hp gate with maximum one-step margin. +- If guided repair is triggered and all four guided candidates are exact + negatives, execute one by the same selector as an offline continuation. It + retains audit truth `y=0` and is stored only in `D0`; it never enters `D+` or + the GP. Collision, success, timeout, and the three-step trap stop remain. +- Update order is whole `D+` and then isolated whole `D0`. Every record is used + exactly once per inner pass. Each population contributes one Adam step per + pass. Canonical baseline: batch 128, lr `1e-5`, one inner pass, alpha 0, + frozen visual encoder, no expert/prox/anchor/curriculum/rollback. +- Every round records fixed-context/fixed-latent before/after probes. Same-lineage + M2 is a fit diagnostic, never a scientific evaluation. + +Claude may add private experimental arms, but must keep a control with these +exact semantics and isolate each algorithmic change. Do not relabel `D0` as a +safety positive, leak it into the GP, or change the locked pretrained/Kazuki +baselines. + +## Evaluation and target + +The four primary metrics are collision rate (lower), window-level Validity +(higher), successful minimum clearance (higher), and successful time-to-goal +(lower, while retaining liveness). Also report SR and timeout. + +Desired gamma family: + +- lower gamma: no higher CR, higher successful clearance, and longer time; +- higher gamma: higher Validity and faster progress. + +Trend eligibility uses adjacent-pair tolerances `rate=.1`, `clearance=.02 m`, +and `time=1 s`; each family must satisfy at least 75% of adjacent pairs. + +Canonical raw evaluation is temperature 1. If temperature tuning is explored, +it must be named explicitly: select a per-gamma schedule on a separate +calibration bank, freeze it before confirmation, and run a fresh disjoint bank. +Never choose temperature after reading confirmation. Report temperature-1 raw +results alongside a tuned deployment result when making a scientific claim. + +The current authenticated fresh disjoint M50 references (350 trajectories per +method, `ep0=485000`) are printed by the start command. They are approximately: + +| method | SR | CR | timeout | Validity | clearance [m] | time [s] | +|---|---:|---:|---:|---:|---:|---:| +| pretrained | .660 | .337 | .003 | .607 | .130 | 8.53 | +| locked Kazuki | .840 | .157 | .003 | .325 | .184 | 4.32 | + +These are cached deployment results, not values to recompute during handoff. + +## Files that define the mechanism + +- `sfm_b1_expert.py`: faithful ID SafeMPPI demonstration expert. +- `stage3_pretrain_sfm.py`: from-scratch Hp10 pretraining and ID-only promotion. +- `grid_policy_sfm.py`, `sfm_hp_history.py`: policy and ten-grid history. +- `sfm_b1_neutral_multiround.py`: canonical macro-round and two-population update. +- `sfm_b1_kazuki_repair_audit.py`: ordinary/guided collection and `D+`/`D0` storage. +- `sfm_b1_rbf.py`, `sfm_b1_store.py`: RBF uncertainty and replay accounting. +- `sfm_metrics2.py`: exact full-H moving-pedestrian certificate implementation. +- `sfm_b1_offline_eval.py`: fixed raw evaluation and four-metric plot contract. +- `run_sfm_neutral_gamma_temperature.py`: separated calibration/confirmation. + +Blind spots to keep visible: same-lineage probes overstate generalization; +temperature schedules can overfit a bank; `D0` is behaviorally useful but +verifier-negative; exact certification and closed-loop collision avoidance are +not the same objective; and longer training has not shown monotonic improvement. diff --git a/overnight_run_07_12_sfm/analysis/test_claude_afe_fidelity.py b/overnight_run_07_12_sfm/analysis/test_claude_afe_fidelity.py new file mode 100644 index 0000000..e0cb3d3 --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_claude_afe_fidelity.py @@ -0,0 +1,112 @@ +import numpy as np +import pytest +import torch + +import claude_afe_fidelity as AF + + +def test_golden_rule_gate(): + ok = dict(resolved=True, y=1, full_h=True, terminal_step=10) + assert AF.certified(ok) + for bad in (dict(ok, y=0), dict(ok, full_h=False), + dict(ok, terminal_step=9), dict(ok, resolved=False)): + assert not AF.certified(bad) + + +def _shard(counts_by_context, round_i=1): + shard = AF.CertifiedQueryShard(round_i) + gammas = (0.1, 0.5, 1.0) + rng = np.random.default_rng(0) + for c, (n_pos, js) in enumerate(counts_by_context): + cid = shard.add_context( + scenario_id=100 + c % 2, gamma=gammas[c % 3], step=c, + state=np.zeros(4, np.float32), + hp10=np.zeros((10, 16, 12), np.float32), + low5=np.zeros(5, np.float32), + hist=np.zeros((16, 2), np.float32), + ped_xy=np.ones((1, 2), np.float32), + ped_vel=np.zeros((1, 2), np.float32), + ) + for j in range(n_pos): + shard.add_query( + cid, rng.uniform(-1, 1, (10, 2)).astype(np.float32), + np.zeros(20, np.float32), + dict(resolved=True, y=1, full_h=True, terminal_step=10), + J=js[j], sigma=0.1, executed=(j == 0), source="pool", + ) + return shard + + +def test_shard_rejects_uncertified_positive(): + shard = AF.CertifiedQueryShard(1) + cid = shard.add_context( + scenario_id=1, gamma=0.5, step=0, state=np.zeros(4, np.float32), + hp10=np.zeros((10, 16, 12), np.float32), + low5=np.zeros(5, np.float32), hist=np.zeros((16, 2), np.float32), + ped_xy=np.ones((1, 2), np.float32), + ped_vel=np.zeros((1, 2), np.float32), + ) + with pytest.raises(ValueError): + shard.add_query( + cid, np.zeros((10, 2), np.float32), np.zeros(20, np.float32), + dict(resolved=True, y=1, full_h=True, terminal_step=9), + J=0.0, sigma=None, executed=False, source="pool", + ) + + +def test_tau_calibration_hits_median_ess_half(): + shard = _shard([(4, [0.0, 1.0, 2.0, 3.0]), (3, [0.0, 0.5, 4.0]), + (5, [0.0, 2.0, 2.5, 3.0, 9.0])]) + tau, ess = AF.calibrate_tau_j(shard) + assert tau is not None + assert ess == pytest.approx(0.5, abs=1e-3) + + +def test_objective_weights_sum_to_one_and_G_is_contextwise_gibbs(): + shard = _shard([(3, [0.0, 1.0, 2.0]), (2, [0.0, 5.0]), (1, [0.0])]) + for objective in ("U", "G"): + records, weights, info = AF.objective_weights([shard], objective) + assert sum(weights.values()) == pytest.approx(1.0, abs=1e-9) + assert info["mass_residual"] < 1e-9 + _, wU, _ = AF.objective_weights([shard], "U") + _, wG, infoG = AF.objective_weights([shard], "G") + rows0 = [r for r in shard.Dplus if r["context_id"] == 0] + u_vals = [wU[(id(shard), r["query_id"])] for r in rows0] + assert max(u_vals) == pytest.approx(min(u_vals)) + g_vals = [(r["J"], wG[(id(shard), r["query_id"])]) for r in rows0] + g_sorted = sorted(g_vals) + assert g_sorted[0][1] > g_sorted[-1][1] + + +class _Tiny(torch.nn.Module): + def __init__(self): + super().__init__() + self.enc_grid = torch.nn.Linear(1, 1, bias=False) + self.head = torch.nn.Linear(20, 20, bias=False) + self.d, self.u_max = 20, 2.0 + + def ctx_from(self, g, l, h): + return l[:, :1] + + def forward(self, v, t, c): + return self.head(v) + + def cfm_loss(self, controls, ctx, weights=None): + per = (self.head(controls.reshape(len(controls), 20) / 2.0) + - controls.reshape(len(controls), 20) / 2.0).square().mean(1) + return (per * weights).mean() if weights is not None else per.mean() + + def module_groups(self): + return {"head": self.head} + + +def test_certified_replay_one_adam_step_per_epoch(): + shard = _shard([(3, [0.0, 1.0, 2.0]), (2, [0.0, 5.0])]) + policy = _Tiny() + policy.enc_grid.weight.requires_grad_(False) + opt = torch.optim.Adam([policy.head.weight], lr=1e-3) + info = AF.certified_replay( + policy, opt, [shard], objective="G", epochs=4, batch=2, + device="cpu", seed=1, + ) + assert info["adam_steps"] == 4 and len(info["losses"]) == 4 diff --git a/overnight_run_07_12_sfm/analysis/test_claude_continuation.py b/overnight_run_07_12_sfm/analysis/test_claude_continuation.py new file mode 100644 index 0000000..9753240 --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_claude_continuation.py @@ -0,0 +1,201 @@ +import numpy as np +import pytest +import torch + +import claude_continuation as CC +import sfm_b1_offline_store as OS + + +def _result(y): + return dict(resolved=True, y=int(y), taskspace=bool(y), + collision_free=bool(y), certificate=bool(y), full_h=True, + terminal_step=10, diagnostics={"m": 1}) + + +def _shard(positive=9, negative=4): + shard = OS.ExecutedRoundShard(2) + gammas = (0.1, 0.2, 0.3, 0.4, 0.5, 0.7, 1.0) + for index in range(positive + negative): + context_id = shard.add_context( + scenario_id=500 + index % 4, gamma=gammas[index % 7], + step=index, state=np.zeros(4, np.float32), + hp10=np.zeros((10, 16, 12), np.float32), + low5=np.zeros(5, np.float32), + hist=np.zeros((16, 2), np.float32), + ped_xy=np.ones((2, 2), np.float32) * 3, + ped_vel=np.zeros((2, 2), np.float32), + ) + shard.add_executed_window( + context_id, np.full((10, 2), 0.1 * index, np.float32), + np.zeros(20, np.float32), _result(index < positive), + execution_source="selected_B", nvp_context=False, + candidate_id=0, acquisition_step=0, sigma=0.1, hp_margin=0.1, + mode="U", + ) + return shard + + +def _anchor(n=6): + records = [] + gammas = (0.1, 0.5, 1.0) + for index in range(n): + records.append(dict( + episode=380000 + index, gamma=gammas[index % 3], step=index, + context=dict( + state=np.zeros(4, np.float32), + hp10=np.zeros((10, 16, 12), np.float32), + low5=np.zeros(5, np.float32), + hist=np.zeros((16, 2), np.float32), + ped_xy=np.ones((2, 2), np.float32) * 3, + ped_vel=np.zeros((2, 2), np.float32), + ), + controls=np.full((10, 2), 0.5, np.float32), + )) + return CC.AnchorShard(dict(records=records)) + + +def test_interleave_covers_exactly_once_and_is_deterministic(): + a = [f"a{i}" for i in range(5)] + b = [f"b{i}" for i in range(2)] + c = [f"c{i}" for i in range(9)] + first = CC._interleave([a, b, c]) + second = CC._interleave([a, b, c]) + assert first == second + assert sorted(first) == sorted(a + b + c) + + +def test_positive_mass_shares_sum_to_declared_composition(): + shard = _shard() + anchor = _anchor() + recovery = [dict( + context_id=0, controls=np.full((10, 2), 0.2, np.float32), y=1, + query_id=0, execution_source="synthetic_certified_recovery", + nvp_context=False, candidate_id=None, acquisition_step=None, + sigma=None, hp_margin=None, mode="recovery", + verifier_diagnostics={}, + )] + for anchor_mass in CC.ANCHOR_GRID: + populations = CC.build_positive_populations( + shard, recovery, anchor, anchor_mass, + ) + weights, membership, total = CC.build_weights( + populations, [(shard, r) for r in shard.Dminus], + ) + sums = {} + for name, _, records, share in populations: + s = sum( + weights[(id(h), int(r["query_id"]))] for h, r in records + ) + sums[name] = s / total + assert s / total == pytest.approx(share, abs=1e-9) + assert sum(sums.values()) == pytest.approx(1.0, abs=1e-9) + assert sums["recovery"] == pytest.approx(0.05, abs=1e-9) + assert sums["anchor"] == pytest.approx(anchor_mass, abs=1e-9) + + +def test_epoch_batches_exact_once_positive_seeded_deterministic(): + shard = _shard() + anchor = _anchor() + populations = CC.build_positive_populations(shard, [], anchor, 0.25) + negatives = [(shard, r) for r in shard.Dminus] + first = CC.epoch_batches( + populations, negatives, batch=4, seed=7, epoch=0, + ) + second = CC.epoch_batches( + populations, negatives, batch=4, seed=7, epoch=0, + ) + def ids(batches): + return [[(id(h), r["query_id"]) for h, r in b] for b in batches] + assert ids(first) == ids(second) + flat = [x for b in ids(first) for x in b] + assert len(flat) == len(set(flat)) + total = sum(len(records) for _, _, records, _ in populations) \ + + len(negatives) + assert len(flat) == total + weights, membership, _ = CC.build_weights(populations, negatives) + for b in first: + assert any( + membership[(id(h), int(r["query_id"]))] != "negative" + for h, r in b + ) + # a different epoch produces a different (but deterministic) order + third = CC.epoch_batches( + populations, negatives, batch=4, seed=7, epoch=1, + ) + assert ids(third) != ids(first) + + +def test_admissibility_hard_gates(): + r1 = dict(SR=0.7, CR=0.3, timeout=0.0, Validity=0.7, clearance=0.12, + time=10.0, g01=dict(clearance=0.15, time=13.0), + g10=dict(clearance=0.11, time=10.0)) + good = dict(r1) + ok, checks = CC.admissible(good, r1) + assert ok, checks + worse_cr = dict(r1, CR=0.31) + assert not CC.admissible(worse_cr, r1)[0] + worse_v = dict(r1, Validity=0.699) + assert not CC.admissible(worse_v, r1)[0] + sr_within_tol = dict(r1, SR=0.66) + assert CC.admissible(sr_within_tol, r1)[0] + sr_beyond = dict(r1, SR=0.64) + assert not CC.admissible(sr_beyond, r1)[0] + bad_trend = dict(r1, g01=dict(clearance=0.10, time=13.0)) + assert not CC.admissible(bad_trend, r1)[0] + no_success_cell = dict(r1, g01=dict(clearance=None, time=None)) + assert not CC.admissible(no_success_cell, r1)[0] + + +class _TinyPolicy(torch.nn.Module): + def __init__(self): + super().__init__() + self.enc_grid = torch.nn.Linear(1, 1, bias=False) + self.head = torch.nn.Linear(20, 20, bias=False) + self.d = 20 + self.u_max = 2.0 + + def ctx_from(self, grid, low, hist): + del grid, hist + return low[:, :1] + + def forward(self, value, tau, context): + del tau, context + return self.head(value) + + def cfm_loss(self, controls, context, weights=None): + del context + value = controls.reshape(len(controls), self.d) / self.u_max + per = (self.head(value) - value).square().mean(dim=1) + if weights is None: + return per.mean() + return (per * weights).sum() / weights.sum() + + def module_groups(self): + return {"E_g": self.enc_grid, "head": self.head} + + +def test_continuation_replay_runs_and_logs_population_losses(): + shard = _shard() + anchor = _anchor() + policy = _TinyPolicy() + for parameter in policy.enc_grid.parameters(): + parameter.requires_grad_(False) + optimizer = torch.optim.Adam( + [p for p in policy.parameters() if p.requires_grad], lr=1e-3, + ) + populations = CC.build_positive_populations(shard, [], anchor, 0.5) + negatives = [(shard, r) for r in shard.Dminus] + before = policy.head.weight.detach().clone() + log = CC.continuation_replay( + policy, optimizer, populations, negatives, epochs=2, batch=4, + seed=3, device="cpu", + ) + assert log["steps"] > 0 + assert log["unique_positive"] == len(shard.Dplus) + len(anchor.windows) + assert log["losses"]["anchor"] is not None + assert log["losses"]["new"] is not None + assert not torch.equal(before, policy.head.weight) + assert torch.equal( + _TinyPolicy().enc_grid.weight * 0 + policy.enc_grid.weight, + policy.enc_grid.weight, + ) diff --git a/overnight_run_07_12_sfm/analysis/test_claude_corrected_distill.py b/overnight_run_07_12_sfm/analysis/test_claude_corrected_distill.py new file mode 100644 index 0000000..e70c9e8 --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_claude_corrected_distill.py @@ -0,0 +1,169 @@ +import json + +import numpy as np +import pytest +import torch + +import claude_corrected_distill as CD +import grid_policy_sfm as GPS +import sfm_b1_offline_store as OS + + +def _result(y): + return dict(resolved=True, y=int(y), taskspace=bool(y), + collision_free=bool(y), certificate=bool(y), full_h=True, + terminal_step=10, diagnostics={"m": 1}) + + +def _shard(round_i, positive=6, negative=2, base=100): + shard = OS.ExecutedRoundShard(round_i) + gammas = (0.1, 0.3, 0.5, 1.0) + rng = np.random.default_rng(round_i) + for index in range(positive + negative): + context_id = shard.add_context( + scenario_id=base + index % 3, gamma=gammas[index % 4], + step=index, state=rng.normal(size=4).astype(np.float32), + hp10=rng.normal(size=(10, 16, 12)).astype(np.float32), + low5=rng.normal(size=5).astype(np.float32), + hist=rng.normal(size=(16, 2)).astype(np.float32), + ped_xy=np.ones((2, 2), np.float32) * 3, + ped_vel=np.zeros((2, 2), np.float32), + ) + shard.add_executed_window( + context_id, + rng.uniform(-1, 1, size=(10, 2)).astype(np.float32), + np.zeros(20, np.float32), _result(index < positive), + execution_source="selected_B", nvp_context=False, + candidate_id=0, acquisition_step=0, sigma=0.1, hp_margin=0.1, + mode="U", + ) + return shard + + +def test_w2_holder_fail_closed(): + with pytest.raises(RuntimeError, match="W=2 requires BOTH"): + CD.CorrectedRecent([_shard(1)]) + with pytest.raises(RuntimeError, match="distinct rounds"): + CD.CorrectedRecent([_shard(1), _shard(1, base=200)]) + recent = CD.CorrectedRecent([_shard(1), _shard(2, base=200)]) + assert len(recent.positive_records()) == 12 + assert len(recent.negative_records()) == 4 + + +def test_production_policy_chunked_gradient_matches_direct(): + torch.manual_seed(0) + policy = GPS.build_sfm_policy(device="cpu") + policy.train() + for parameter in policy.enc_grid.parameters(): + parameter.requires_grad_(False) + shard = _shard(3, positive=9, negative=0) + records = [(shard, row) for row in shard.Dplus] + mass, _ = CD.normalized_teacher_mass(records) + batch, seed = 4, 777 + + # production chunked-accumulated path + policy.zero_grad(set_to_none=True) + chunked_total = CD.teacher_epoch_loss( + policy, records, mass, batch=batch, device="cpu", seed=seed, + backward=True, + ) + chunked_gradients = { + name: parameter.grad.detach().clone() + for name, parameter in policy.named_parameters() + if parameter.requires_grad and parameter.grad is not None + } + + # independent direct implementation replaying the identical chunk noise + policy.zero_grad(set_to_none=True) + direct = CD.direct_teacher_loss( + policy, records, mass, batch=batch, device="cpu", seed=seed, + ) + direct.backward() + direct_gradients = { + name: parameter.grad.detach().clone() + for name, parameter in policy.named_parameters() + if parameter.requires_grad and parameter.grad is not None + } + + assert chunked_total == pytest.approx(float(direct.detach()), abs=1e-6) + assert set(chunked_gradients) == set(direct_gradients) + worst = max( + float((chunked_gradients[name] - direct_gradients[name]).abs().max()) + for name in chunked_gradients + ) + assert worst < 1e-5, f"gradient disagreement {worst}" + # the objective actually moves parameters (non-trivial gradients) + assert max( + float(g.abs().max()) for g in chunked_gradients.values() + ) > 0.0 + + +class _TinyPolicy(torch.nn.Module): + def __init__(self): + super().__init__() + self.enc_grid = torch.nn.Linear(1, 1, bias=False) + self.head = torch.nn.Linear(20, 20, bias=False) + self.d = 20 + self.u_max = 2.0 + + def ctx_from(self, grid, low, hist): + del grid, hist + return low[:, :1] + + def forward(self, value, tau, context): + del tau, context + return self.head(value) + + def cfm_loss(self, controls, context, weights=None): + del context + value = controls.reshape(len(controls), self.d) / self.u_max + per = (self.head(value) - value).square().mean(dim=1) + if weights is not None: + per = per * weights + return per.mean() + + def module_groups(self): + return {"E_g": self.enc_grid, "head": self.head} + + +def test_corrected_ordinary_is_one_adam_step_per_epoch(): + recent = CD.CorrectedRecent([_shard(1), _shard(2, base=200)]) + policy = _TinyPolicy() + for parameter in policy.enc_grid.parameters(): + parameter.requires_grad_(False) + optimizer = torch.optim.Adam( + [p for p in policy.parameters() if p.requires_grad], lr=1e-3, + ) + out = CD.corrected_ordinary_epochs( + policy, optimizer, recent, alpha=0.01, epochs=3, batch=4, + device="cpu", seed=9, + ) + assert out["adam_steps"] == 3 + assert all(row["optimizer_steps"] == 1 for row in out["per_epoch"]) + + +def test_pooled_by_label_reads_the_right_record(tmp_path): + def _cell(sr): + return dict(summary=dict( + pooled=dict( + SR=sr, CR=0.2, timeout=0.0, + Validity=dict(mean=0.7), + successful_clearance=dict(mean=0.1), + successful_time_to_goal=dict(mean=9.0), + ), + per_gamma={ + g: dict(successful_clearance=dict(mean=0.1), + successful_time_to_goal=dict(mean=9.0)) + for g in ("0.1", "1.0") + }, + )) + payload = dict(records=[ + dict(label="r0", cell=_cell(0.5)), + dict(label="r1", cell=_cell(0.9)), + ]) + path = tmp_path / "metrics.json" + path.write_text(json.dumps(payload)) + assert CD.pooled_by_label(path, "r1")["SR"] == 0.9 + assert CD.pooled_by_label(path, "r0")["SR"] == 0.5 + with pytest.raises(KeyError): + CD.pooled_by_label(path, "r7") diff --git a/overnight_run_07_12_sfm/analysis/test_claude_mpc_pool.py b/overnight_run_07_12_sfm/analysis/test_claude_mpc_pool.py new file mode 100644 index 0000000..2f737ec --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_claude_mpc_pool.py @@ -0,0 +1,129 @@ +import numpy as np +import torch + +import claude_mpc_pool as MP +import sfm_b1_offline_store as OS +import sfm_scene as SS + + +def _result_positive(): + return dict( + resolved=True, y=1, taskspace=True, collision_free=True, + certificate=True, full_h=True, terminal_step=10, + diagnostics={"slack": 0.1}, + ) + + +def _record_episode_shard(scenario=123, gamma=0.5, steps=4, n_ped=6): + """Simulate a short episode and store its exact contexts/windows.""" + speed_range = (1.0, 2.0) + humans = SS.make_humans(scenario, 0, n_ped, speed_range) + state = np.zeros(4, np.float32) + shard = OS.ExecutedRoundShard(1) + rng = np.random.default_rng(0) + for step in range(steps): + ped_xy, ped_vel = SS.collect_humans(humans) + context_id = shard.add_context( + scenario_id=scenario, gamma=gamma, step=step, state=state.copy(), + hp10=np.zeros((10, 16, 12), np.float32), + low5=np.zeros(5, np.float32), + hist=np.zeros((16, 2), np.float32), + ped_xy=ped_xy.copy(), ped_vel=ped_vel.copy(), + ) + controls = np.tile( + rng.uniform(-0.5, 0.5, size=(1, 2)).astype(np.float32), (10, 1), + ) + shard.add_executed_window( + context_id, controls, np.zeros(20, np.float32), + _result_positive(), execution_source="selected_B", + nvp_context=False, candidate_id=0, acquisition_step=0, + sigma=0.1, hp_margin=0.1, mode="U", + ) + action = controls[0] + state[:2] = state[:2] + SS.DT * state[2:4] + 0.5 * SS.DT ** 2 * action + state[2:4] = state[2:4] + SS.DT * action + SS.advance_humans(humans, state) + return shard + + +def test_prefix_replay_reconstructs_stored_context_exactly(): + shard = _record_episode_shard() + environment = dict(n_ped=6, ped_speed_range=(1.0, 2.0)) + for target_step in (0, 2, 3): + context = next( + c for c in shard.contexts if int(c["step"]) == target_step + ) + humans, state = MP.replay_prefix_humans( + shard, 123, 0.5, target_step, environment, + ) + ped_xy, ped_vel = SS.collect_humans(humans) + assert np.allclose(ped_xy, context["ped_xy"], atol=1e-6) + assert np.allclose(ped_vel, context["ped_vel"], atol=1e-6) + assert np.allclose(state, context["state"], atol=1e-6) + + +def test_privileged_config_matches_codex_source(): + cfg = MP.privileged_sfm_config() + assert cfg.exact_sfm_step_filter is True + assert cfg.step_filter_margin == 0.22 + assert cfg.step_filter_goal_plans == 12 + assert cfg.step_filter_avoid_plans == 18 + assert cfg.safe_coef_by_gamma == (1.0, 0.3, 1.0, 0.3, 0.3, 0.3, 0.1) + + +class _TinyPolicy(torch.nn.Module): + def __init__(self): + super().__init__() + self.enc_grid = torch.nn.Linear(1, 1, bias=False) + self.head = torch.nn.Linear(20, 20, bias=False) + self.d = 20 + self.u_max = 2.0 + + def ctx_from(self, grid, low, hist): + del grid, hist + return low[:, :1] + + def forward(self, value, tau, context): + del tau, context + return self.head(value) + + def cfm_loss(self, controls, context, weights=None): + del context + value = controls.reshape(len(controls), self.d) / self.u_max + per = (self.head(value) - value).square().mean(dim=1) + if weights is None: + return per.mean() + return (per * weights).sum() / weights.sum() + + +def test_distill_block_trains_only_on_mpc_records_and_moves_params(): + shard = _record_episode_shard(steps=3) + records = [ + dict( + context_id=int(c["context_id"]), y=1, query_id=i, + controls=np.full((10, 2), 0.3, np.float32), + source="codex_privileged_mpc_pool", rank=0, + privileged_clearance=0.3, privileged_margin=0.22, + safemppi_cost=1.0, verifier_diagnostics={}, + ) + for i, c in enumerate(shard.contexts) + ] + policy = _TinyPolicy() + for parameter in policy.enc_grid.parameters(): + parameter.requires_grad_(False) + optimizer = torch.optim.Adam( + [p for p in policy.parameters() if p.requires_grad], lr=1e-3, + ) + before = {k: v.clone() for k, v in policy.state_dict().items()} + result = MP.distill_block( + policy, optimizer, shard, records, epochs=2, batch=2, seed=5, + ) + assert result["steps"] == 2 * 2 # ceil(3/2)=2 batches x 2 epochs + assert result["records"] == 3 + assert not torch.equal(before["head.weight"], policy.state_dict()["head.weight"]) + assert torch.equal(before["enc_grid.weight"], policy.state_dict()["enc_grid.weight"]) + # empty buffer is a no-op + empty = MP.distill_block( + policy, optimizer, shard, [], epochs=2, batch=2, seed=5, + ) + assert empty["steps"] == 0 diff --git a/overnight_run_07_12_sfm/analysis/test_claude_offline_aug.py b/overnight_run_07_12_sfm/analysis/test_claude_offline_aug.py new file mode 100644 index 0000000..4a8c75c --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_claude_offline_aug.py @@ -0,0 +1,292 @@ +import copy + +import numpy as np +import torch + +import claude_offline_aug as AUG +import sfm_b1_offline_replay as OR +import sfm_b1_offline_store as OS +import sfm_metrics2 as SM + + +def _result(y): + return dict( + resolved=True, y=int(y), taskspace=bool(y), collision_free=bool(y), + certificate=bool(y), full_h=True, terminal_step=10, + diagnostics={"margin": 0.25}, + ) + + +def _context(shard, *, scenario, gamma, step, state, ped_xy, ped_vel): + return shard.add_context( + scenario_id=scenario, gamma=gamma, step=step, + state=np.asarray(state, np.float32), + hp10=np.zeros((10, 16, 12), np.float32), + low5=np.zeros(5, np.float32), + hist=np.zeros((16, 2), np.float32), + ped_xy=np.asarray(ped_xy, np.float32).reshape(-1, 2), + ped_vel=np.asarray(ped_vel, np.float32).reshape(-1, 2), + ) + + +def _add(shard, *, scenario, gamma, step, y, state, ped_xy, ped_vel, + controls, collision_after=False, trap=False, result=None): + context_id = _context( + shard, scenario=scenario, gamma=gamma, step=step, state=state, + ped_xy=ped_xy, ped_vel=ped_vel, + ) + window_id = shard.add_executed_window( + context_id, np.asarray(controls, np.float32), + np.zeros(20, np.float32), result or _result(y), + execution_source="selected_B" if y else "raw_continuation", + nvp_context=not bool(y), candidate_id=0 if y else None, + acquisition_step=0 if y else None, sigma=0.4 if y else None, + hp_margin=0.2, mode="U", + ) + shard.windows[window_id].update( + collision_after_action=bool(collision_after), trap_event=bool(trap), + ) + return window_id + + +FAST = np.full((10, 2), 1.2, np.float32) # strong forward motion +STILL = np.zeros((10, 2), np.float32) # no displacement +FAR_PED = [[5.5, 0.5]] +NEAR_PED = [[0.35, 0.12]] +ZERO_VEL = [[0.0, 0.0]] + + +class _TinyPolicy(torch.nn.Module): + def __init__(self): + super().__init__() + self.enc_grid = torch.nn.Linear(1, 1, bias=False) + self.head = torch.nn.Linear(20, 20, bias=False) + self.d = 20 + self.u_max = 2.0 + + def ctx_from(self, grid, low, hist): + del grid, hist + return low[:, :1] + + def forward(self, value, tau, context): + del tau, context + return self.head(value) + + def cfm_loss(self, controls, context, weights=None): + del context + value = controls.reshape(len(controls), self.d) / self.u_max + per = (self.head(value) - value).square().mean(dim=1) + if weights is None: + return per.mean() + return (per * weights).sum() / weights.sum() + + def module_groups(self): + return {"E_g": self.enc_grid, "head": self.head} + + +def _freeze_encoder(policy): + for parameter in policy.enc_grid.parameters(): + parameter.requires_grad_(False) + + +def _mixed_shard(): + shard = OS.ExecutedRoundShard(1) + gammas = (0.1, 0.2, 0.3, 0.4, 0.5, 0.7, 1.0) + for index in range(10): + _add( + shard, scenario=100 + index, gamma=gammas[index % 7], step=index, + y=index < 7, state=[0.5, 0.5, 0.5, 0.5], + ped_xy=NEAR_PED if index % 2 else FAR_PED, ped_vel=ZERO_VEL, + controls=FAST if index < 7 else STILL, + collision_after=(index == 8), trap=(index == 9), + ) + return shard + + +class _InlineExecutor: + def map(self, fn, tasks): + return [fn(task) for task in tasks] + + +def test_original_mode_returns_untouched_shard(): + shard = _mixed_shard() + view, report = AUG.build_replay_view(shard, "original") + assert view is shard + assert report["mode"] == "original" + + +def test_original_mode_replay_is_bitwise_identical_to_control(): + shard = _mixed_shard() + torch.manual_seed(0) + policy_a = _TinyPolicy() + policy_b = copy.deepcopy(policy_a) + _freeze_encoder(policy_a) + _freeze_encoder(policy_b) + opt_a = torch.optim.Adam( + [p for p in policy_a.parameters() if p.requires_grad], lr=1e-3, + ) + opt_b = torch.optim.Adam( + [p for p in policy_b.parameters() if p.requires_grad], lr=1e-3, + ) + control = OR.replay( + policy_a, opt_a, shard, alpha=0.01, exposure_epochs=1, batch=4, + device="cpu", seed=7, + ) + treated = AUG.replay_with_mode( + policy_b, opt_b, shard, mode="original", alpha=0.01, + exposure_epochs=1, batch=4, device="cpu", seed=7, + ) + for key, value in policy_a.state_dict().items(): + assert torch.equal(value, policy_b.state_dict()[key]), key + assert control["optimizer_steps"] == treated["optimizer_steps"] + + +def test_population_tagging_follows_declared_rules(): + shard = OS.ExecutedRoundShard(2) + # near + moving certified positive -> pop A + _add(shard, scenario=1, gamma=0.5, step=0, y=1, + state=[0, 0, 1.0, 0], ped_xy=[[0.8, 0.6]], ped_vel=ZERO_VEL, + controls=FAST) + # far certified positive -> excluded from pop A + _add(shard, scenario=2, gamma=0.5, step=0, y=1, + state=[0, 0, 1.0, 0], ped_xy=FAR_PED, ped_vel=ZERO_VEL, + controls=np.full((10, 2), 0.4, np.float32)) + # certified but still (slow) positive -> excluded from A, NEVER in B + _add(shard, scenario=3, gamma=0.5, step=0, y=1, + state=[0, 0, 0, 0], ped_xy=[[0.8, 0.6]], ped_vel=ZERO_VEL, + controls=STILL) + # negative with actual collision -> pop B + _add(shard, scenario=4, gamma=0.5, step=0, y=0, + state=[0, 0, 1.0, 0], ped_xy=NEAR_PED, ped_vel=ZERO_VEL, + controls=FAST, collision_after=True) + # negative, no collision flags, but no displacement -> pop B (no progress) + _add(shard, scenario=5, gamma=0.5, step=0, y=0, + state=[0, 0, 0, 0], ped_xy=FAR_PED, ped_vel=ZERO_VEL, + controls=STILL) + pop_a, pop_b, stats = AUG.tag_populations(shard) + assert [w["context_id"] for w in pop_a] == [0] + assert sorted(w["context_id"] for w in pop_b) == [3, 4] + assert stats["popA"] == 1 and stats["popB"] == 2 + # the certified-slow window is not relabeled + assert all(int(w["y"]) == 0 for w in pop_b) + + +def test_recovery_candidates_are_deterministic_and_bounded(): + state = [1.0, 1.0, 1.5, -0.5] + first = AUG.recovery_candidates(state) + second = AUG.recovery_candidates(state) + assert len(first) == 1 + len(AUG.BRAKE_STEPS) * AUG.K_DIR * len(AUG.ACCELS) + for (controls_a, prov_a), (controls_b, prov_b) in zip(first, second): + assert np.array_equal(controls_a, controls_b) + assert prov_a == prov_b + assert controls_a.shape == (10, 2) + assert float(np.abs(controls_a).max()) <= 2.0 + 1e-6 + + +def test_recovery_records_are_exactly_certified_with_provenance(): + shard = OS.ExecutedRoundShard(3) + _add(shard, scenario=9, gamma=0.5, step=4, y=0, + state=[2.0, 2.0, 0.8, 0.0], ped_xy=[[2.9, 2.0]], + ped_vel=[[-0.5, 0.0]], controls=FAST, collision_after=True) + records, audit = AUG.build_recovery_records( + shard, shard.Dminus, _InlineExecutor(), + ) + assert records, "an escapable context must yield certified recovery" + assert len(records) <= AUG.RECOVERY_KEEP + assert audit["certified_kept"] == len(records) + for record in records: + context = shard.contexts[record["context_id"]] + recheck = SM.verify_query( + context["state"], record["controls"], context["ped_xy"], + context["ped_vel"], context["gamma"], + ) + assert recheck["resolved"] and int(recheck["y"]) == 1 + assert record["execution_source"] == "synthetic_certified_recovery" + prov = record["recovery_provenance"] + assert prov["parent_context_id"] == record["context_id"] + assert "generator" in prov and "objective" in prov + assert prov["family"] == "v1" + + +def test_orig_plus_recovery_keeps_full_population_and_appends(): + shard = _mixed_shard() + view, report = AUG.build_replay_view( + shard, "orig_plus_recovery", executor=_InlineExecutor(), + ) + added = report["recovery_audit"]["certified_kept"] + assert len(view.windows) == len(shard.windows) + added + assert len(view.Dminus) == len(shard.Dminus) + assert len(view.Dplus) == len(shard.Dplus) + added + synthetic = [ + w for w in view.windows + if w["execution_source"] == "synthetic_certified_recovery" + ] + assert len(synthetic) == added + + +def test_v2_family_is_deterministic_goal_directed_and_certifiable(): + state = [1.0, 1.0, 0.4, -0.6] + first = AUG.recovery_candidates_v2(state) + second = AUG.recovery_candidates_v2(state) + assert len(first) == len(second) == ( + len(AUG.CRUISE_SPEEDS) + + 2 * AUG.K_DIR * len(AUG.V2_ACCELS) * len(AUG.CRUISE_SPEEDS) + ) + import sfm_metrics2 as SM2 + for (ca, pa), (cb, pb) in zip(first, second): + assert np.array_equal(ca, cb) and pa == pb + assert ca.shape == (10, 2) + assert float(np.abs(ca).max()) <= 2.0 + 1e-6 + # pure-cruise candidate ends moving toward the goal + controls = first[0][0] + seg = SM2.rollout_positions(state, controls) + velocity = np.asarray(state, np.float32)[2:4].copy() + for action in controls: + velocity = velocity + 0.1 * action + toward = velocity @ ((np.array([6.0, 6.0]) - seg[-1]) + / np.linalg.norm(np.array([6.0, 6.0]) - seg[-1])) + assert toward > 0.5 + + +def test_v2_recovery_records_certified_and_tagged(): + shard = OS.ExecutedRoundShard(4) + # pedestrian behind the robot, receding: a goal-directed certified + # escape clearly exists + _add(shard, scenario=9, gamma=0.5, step=4, y=0, + state=[2.0, 2.0, 0.6, 0.6], ped_xy=[[1.2, 2.0]], + ped_vel=[[-0.6, 0.0]], controls=STILL, trap=True) + records, audit = AUG.build_recovery_records( + shard, shard.Dminus, _InlineExecutor(), family="v2", + ) + assert audit["family"] == "v2" + assert records + for record in records: + assert record["recovery_provenance"]["family"] == "v2" + context = shard.contexts[record["context_id"]] + recheck = SM.verify_query( + context["state"], record["controls"], context["ped_xy"], + context["ped_vel"], context["gamma"], + ) + assert recheck["resolved"] and int(recheck["y"]) == 1 + + +def test_hard_recovery_replay_respects_exact_once_accounting(): + shard = _mixed_shard() + policy = _TinyPolicy() + _freeze_encoder(policy) + optimizer = torch.optim.Adam( + [p for p in policy.parameters() if p.requires_grad], lr=1e-3, + ) + result = AUG.replay_with_mode( + policy, optimizer, shard, mode="hard_recovery", alpha=0.01, + exposure_epochs=1, batch=4, device="cpu", seed=11, + executor=_InlineExecutor(), + ) + report = result["replay_intervention"] + assert report["mode"] == "hard_recovery" + assert "recovery_audit" in report + view = report["view"] + assert view["D"] == view["Dplus"] + view["Dminus"] + assert result["positive_eligible"] == view["Dplus"] + assert result["negative_eligible"] == view["Dminus"] + assert result["exact_once_per_exposure_epoch"] is True diff --git a/overnight_run_07_12_sfm/analysis/test_claude_unverified_mpc_teacher.py b/overnight_run_07_12_sfm/analysis/test_claude_unverified_mpc_teacher.py new file mode 100644 index 0000000..89871a7 --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_claude_unverified_mpc_teacher.py @@ -0,0 +1,262 @@ +import os + +import numpy as np +import pytest +import torch + +import claude_unverified_mpc_teacher as T +import claude_mpc_pool as MP +import sfm_b1_offline_store as OS +import sfm_scene as SS + + +def _result(y): + return dict( + resolved=True, + y=int(y), + taskspace=bool(y), + collision_free=bool(y), + certificate=bool(y), + full_h=True, + terminal_step=10, + diagnostics={"slack": 0.1}, + ) + + +def _shard(round_i=1): + shard = OS.ExecutedRoundShard(round_i) + for scenario, gamma in ((10, 0.1), (11, 0.1), (12, 1.0), (13, 1.0)): + context_id = shard.add_context( + scenario_id=scenario, + gamma=gamma, + step=0, + state=np.zeros(4, np.float32), + hp10=np.zeros((10, 16, 12), np.float32), + low5=np.zeros(5, np.float32), + hist=np.zeros((16, 2), np.float32), + ped_xy=np.array([[2.0, 2.0]], np.float32), + ped_vel=np.zeros((1, 2), np.float32), + ) + shard.add_executed_window( + context_id, + np.zeros((10, 2), np.float32), + np.zeros(20, np.float32), + _result(context_id % 2), + execution_source="unit_test", + nvp_context=False, + ) + return shard + + +class _TinyPolicy(torch.nn.Module): + def __init__(self): + super().__init__() + self.head = torch.nn.Linear(20, 20, bias=False) + self.d = 20 + self.u_max = 2.0 + + def ctx_from(self, hp10, low, hist): + del hp10, hist + return low[:, :1] + + def cfm_loss(self, controls, context, weights=None): + del context + value = controls.reshape(len(controls), 20) / self.u_max + per = (self.head(value) - value).square().mean(dim=1) + if weights is None: + return per.mean() + return (per * weights).sum() / weights.sum() + + +class _InlineExecutor: + def map(self, function, tasks): + return [function(task) for task in tasks] + + +def _teacher_records(shard): + records = [] + for teacher_id, context in enumerate(shard.contexts): + records.append({ + "teacher_id": teacher_id, + "round": shard.round_i, + "context_id": context["context_id"], + "scenario_id": context["scenario_id"], + "gamma": context["gamma"], + "episode_id": context["scenario_id"], + "step": context["step"], + "controls": np.full((10, 2), 0.25, np.float32), + "source": T.TEACHER_SOURCE, + "candidate_source": { + "family": "constant_acceleration", + "source_index": 0, + }, + "controller_config_hash": "abc", + "selector_diagnostics": {"filter_feasible": True}, + "pool_manifest": {}, + "clip": { + "u_max": SS.U_MAX, + "max_abs_before": 0.25, + "max_clip_delta": 0.0, + }, + "context_snapshot": T._context_snapshot(context), + "socp_audit": None, + }) + return records + + +def test_control_contract_rejects_shape_nan_and_large_clipping(): + controls, audit = T.validate_controls( + np.full((10, 2), SS.U_MAX, np.float32), + ) + assert controls.shape == (10, 2) + assert audit["max_clip_delta"] == 0.0 + with pytest.raises(ValueError, match="shape"): + T.validate_controls(np.zeros((9, 2), np.float32)) + bad = np.zeros((10, 2), np.float32) + bad[0, 0] = np.nan + with pytest.raises(ValueError, match="finite"): + T.validate_controls(bad) + too_large = np.zeros((10, 2), np.float32) + too_large[0, 0] = SS.U_MAX + 0.01 + with pytest.raises(ValueError, match="silently reinterpret"): + T.validate_controls(too_large) + + +def test_candidate_provenance_keeps_first_family_when_plans_deduplicate(): + state = np.zeros(4, np.float32) + plan = np.zeros((10, 2), np.float32) + unique, sources = MP._deduplicate_plans_with_sources([ + (plan, {"family": "nominal", "source_index": 0, "state": state}), + (plan.copy(), {"family": "refined", "source_index": 3, "state": state}), + ]) + assert len(unique) == 1 + assert sources == [{"family": "nominal", "source_index": 0}] + + +def test_balanced_context_selection_and_hierarchical_mass(): + shard = _shard() + chosen = T._balanced_context_ids(shard, shard.windows, max_contexts=2) + assert len(chosen) == 2 + assert { + round(float(shard.contexts[index]["gamma"]), 8) for index in chosen + } == {0.1, 1.0} + records = _teacher_records(shard) + mass, accounting = T._teacher_mass(shard, records) + assert accounting["total"] == pytest.approx(1.0) + assert accounting["gamma"]["0.1"] == pytest.approx(0.5) + assert accounting["gamma"]["1"] == pytest.approx(0.5) + assert set(mass) == {0, 1, 2, 3} + + +def test_buffer_never_duck_types_ordinary_D_and_authenticates_shard(tmp_path): + shard = _shard() + shard_path = os.fspath(tmp_path / "round.pt") + buffer_path = os.fspath(tmp_path / "teacher.pt") + shard.save(shard_path) + records = _teacher_records(shard) + payload = T.save_buffer( + buffer_path, shard_path, shard, records, {"counts": {"kept": 4}}, + ) + assert payload["status"] == T.BUFFER_STATUS + assert payload["round_shard_sha256"] == OS.sha256_file(shard_path) + forbidden = {"y", "query_id", "train_eligible", "x0"} + assert all(not (forbidden & set(row)) for row in payload["records"]) + + +def test_socp_negative_is_audit_only_and_teacher_is_still_kept(monkeypatch): + shard = _shard() + + def fake_replay(*args, **kwargs): + del args, kwargs + return [object()], np.zeros(4, np.float32) + + monkeypatch.setattr(T.MP, "replay_prefix_humans", fake_replay) + monkeypatch.setattr( + T.SS, + "collect_humans", + lambda humans: ( + np.array([[2.0, 2.0]], np.float32), + np.zeros((1, 2), np.float32), + ), + ) + monkeypatch.setattr( + T, + "select_teacher_plan", + lambda *args, **kwargs: ( + np.full((10, 2), 0.2, np.float32), + { + "candidate_source": { + "family": "avoidance", + "source_index": 2, + }, + "selector_diagnostics": {"filter_feasible": True}, + "pool_manifest": {}, + "clip": { + "u_max": SS.U_MAX, + "max_abs_before": 0.2, + "max_clip_delta": 0.0, + }, + }, + ), + ) + monkeypatch.setattr( + T.SM, + "verify_in_worker", + lambda task: ( + task[0], + task[1], + { + "resolved": True, + "y": 0, + "full_h": True, + "diagnostics": {"slack": -0.2}, + }, + ), + ) + records, audit = T.harvest_round( + object(), + shard, + [shard.windows[0]], + device="cpu", + environment={"n_ped": 1, "ped_speed_range": (1.0, 2.0)}, + executor=_InlineExecutor(), + audit_socp=True, + ) + assert len(records) == 1 + assert records[0]["socp_audit"]["verifier_label"] == 0 + assert "y" not in records[0] + assert audit["counts"]["kept"] == 1 + assert audit["counts"]["socp_audit_nonpositive"] == 1 + + +def test_teacher_block_is_separate_and_uses_a_fresh_seed_per_step(): + shard = _shard() + records = _teacher_records(shard) + policy = _TinyPolicy() + optimizer = torch.optim.Adam(policy.parameters(), lr=1.0e-3) + before = policy.head.weight.detach().clone() + result = T.distill_block( + policy, + optimizer, + shard, + records, + epochs=2, + batch=2, + seed=7, + ) + assert result["steps"] == 2 + assert result["optimizer_steps"] == 2 + assert result["sample_exposures"] == 8 + assert result["fresh_base_seed_count"] == 4 + assert not torch.equal(before, policy.head.weight) + assert all("x0" not in record for record in records) + noop = T.distill_block( + policy, + optimizer, + shard, + records, + epochs=0, + batch=2, + seed=7, + ) + assert noop["steps"] == 0 diff --git a/overnight_run_07_12_sfm/analysis/test_run_sfm_b1_offline_9arm.py b/overnight_run_07_12_sfm/analysis/test_run_sfm_b1_offline_9arm.py new file mode 100644 index 0000000..f086e51 --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_run_sfm_b1_offline_9arm.py @@ -0,0 +1,193 @@ +from __future__ import annotations + +import json +from pathlib import Path +import sys +from types import SimpleNamespace + +import pytest + + +HERE = Path(__file__).resolve().parents[1] +if str(HERE) not in sys.path: + sys.path.insert(0, str(HERE)) + +import run_sfm_b1_offline_9arm as L # noqa: E402 +import recover_sfm_b1_offline_evaluation as R # noqa: E402 + + +def _gpu(index: int) -> L.BASE.GPU: + return L.BASE.GPU( + index=str(index), + uuid=f"GPU-{index}", + name="test", + memory_total_mib=100, + memory_used_mib=0, + utilization_percent=0, + pci_bus_id=f"0000:0{index}:00.0", + ) + + +def test_arm_grid_and_two_gpu_allocation(): + arms = list(L.arm_grid()) + assert len(arms) == 9 + assert len({arm.name for arm in arms}) == 9 + assert { + (arm.alpha, arm.exposure_epochs) for arm in arms + } == { + (alpha, epochs) + for alpha in L.ALPHAS + for epochs in L.EXPOSURE_EPOCHS + } + allocation = L.allocate_arms(arms, [_gpu(1), _gpu(3)]) + assert sorted(map(len, allocation.values())) == [4, 5] + assert set().union(*map(set, allocation.values())) == set(arms) + assert {arm.exposure_epochs for arm in allocation["GPU-1"]} == { + 1, 10, 100, + } + cost_arms = list(L.arm_grid("safemppi_cost")) + assert len(cost_arms) == 9 + assert all("safemppi_cost" in arm.name for arm in cost_arms) + assert all(arm.selector == "safemppi_cost" for arm in cost_arms) + balanced = list(L.arm_grid("balanced_rank")) + assert len(balanced) == 9 + assert all("balanced_rank" in arm.name for arm in balanced) + + +def test_output_root_must_be_new_and_under_research1(tmp_path, monkeypatch): + root = tmp_path / "research1" + root.mkdir() + monkeypatch.setattr(L, "RESEARCH_ROOT", root) + target = root / "new-study" + assert L._validated_output_root(target) == target.resolve() + target.mkdir() + with pytest.raises(FileExistsError): + L._validated_output_root(target) + with pytest.raises(ValueError): + L._validated_output_root(tmp_path / "elsewhere") + + +def test_commands_cover_declared_rounds_and_raw_common_bank(tmp_path): + checkpoint = tmp_path / "checkpoint.pt" + checkpoint.write_bytes(b"x") + args = type("Args", (), { + "checkpoint": str(checkpoint), + "verifier_workers": 8, + "seed": 20260724, + "eval_ep0": 260000, + "eval_noise_seed": 20260723, + })() + arm = L.Arm(0.01, 10) + train = L._trainer_command(args, arm, tmp_path / "train") + assert train[train.index("--rounds") + 1] == "10" + assert train[train.index("--exposure-epochs") + 1] == "10" + assert train[train.index("--selector") + 1] == "margin" + cost_train = L._trainer_command( + args, L.Arm(0.01, 10, "safemppi_cost"), tmp_path / "cost", + ) + assert cost_train[cost_train.index("--selector") + 1] == "safemppi_cost" + evaluate = L._evaluation_command( + args, + arm, + tmp_path / "train", + tmp_path / "eval", + cache_dir=tmp_path / "common_cache", + ) + checkpoints_start = evaluate.index("--checkpoints") + 1 + checkpoints_end = evaluate.index("--labels") + assert evaluate[checkpoints_start] == str(checkpoint.resolve()) + assert evaluate[checkpoints_start + 1].endswith("round_01.pt") + labels_start = evaluate.index("--labels") + 1 + labels_end = evaluate.index("--scene-profile") + assert evaluate[labels_start:labels_end] == [ + f"r{round_i}" for round_i in range(11) + ] + assert evaluate[evaluate.index("--ep0") + 1] == "260000" + assert evaluate[evaluate.index("--device") + 1] == "cuda:0" + assert evaluate[evaluate.index("--cache-dir") + 1] == str( + (tmp_path / "common_cache").resolve() + ) + common = L._common_r0_command(args, tmp_path / "common") + assert common[common.index("--labels") + 1] == "r0" + assert common[common.index("--checkpoints") + 1] == str( + checkpoint.resolve() + ) + + +def test_screening_key_is_safety_first(): + base = { + "CR": 0.1, + "Validity": 0.5, + "SR": 0.8, + "clearance": 0.1, + "time_to_goal": 9.0, + "round": 1, + "exposure_epochs": 1, + "alpha": 0.0, + } + lower_collision = {**base, "CR": 0.09, "Validity": 0.0} + higher_validity = {**base, "Validity": 0.6, "SR": 0.0} + assert L._screening_key(lower_collision) < L._screening_key(base) + assert L._screening_key(higher_validity) < L._screening_key(base) + + +def test_validate_sidecar_authenticates_digest(tmp_path): + artifact = tmp_path / "round_00.pt" + artifact.write_bytes(b"checkpoint") + sidecar = Path(str(artifact) + ".COMPLETE.json") + sidecar.write_text(json.dumps({ + "status": "COMPLETE", + "sha256": L.BASE.sha256_file(artifact), + })) + observed = L._validate_sidecar(artifact) + assert observed["sha256"] == L.BASE.sha256_file(artifact) + sidecar.write_text(json.dumps({ + "status": "COMPLETE", + "sha256": "0" * 64, + })) + with pytest.raises(RuntimeError): + L._validate_sidecar(artifact) + + +def test_launch_pending_returns_generated_log_path(tmp_path): + job = { + "arm": L.PhaseName("common_r0"), + "gpu": _gpu(1), + "cpu_pool": [0], + "command": [sys.executable, "-c", "print('ok')"], + } + logs = L.BASE._launch_pending([job], tmp_path / "logs") + assert logs == [str((tmp_path / "logs" / "common_r0.log").resolve())] + assert Path(logs[0]).read_text().strip() == "ok" + + +def test_recovery_refuses_delivery_overwrite(tmp_path, monkeypatch): + root = tmp_path / "research1" + run_root = root / "run" + run_root.mkdir(parents=True) + (run_root / "DELIVERY_COMPLETE.json").write_text("{}") + monkeypatch.setattr(L, "RESEARCH_ROOT", root) + with pytest.raises(FileExistsError): + R.recover(SimpleNamespace( + run_root=str(run_root), + gpu_indices="1,3", + idle_memory_mib=1024, + idle_utilization_percent=5, + )) + + +def test_legacy_margin_recipe_normalization_is_margin_only(): + recipe = {"alpha": 0.0} + assert L._normalized_training_recipe(recipe, "margin") == { + "alpha": 0.0, + "selector": "margin", + } + assert L._normalized_training_recipe(recipe, "safemppi_cost") == recipe + + +def test_recovery_allocation_supports_one_or_two_idle_gpus(): + arms = list(L.arm_grid("balanced_rank")) + one = R._recovery_allocation(arms, [_gpu(1)]) + assert one == {"GPU-1": arms} + two = R._recovery_allocation(arms, [_gpu(1), _gpu(3)]) + assert sorted(map(len, two.values())) == [4, 5] diff --git a/overnight_run_07_12_sfm/analysis/test_run_sfm_b1_r2_9arm.py b/overnight_run_07_12_sfm/analysis/test_run_sfm_b1_r2_9arm.py new file mode 100644 index 0000000..5b4951d --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_run_sfm_b1_r2_9arm.py @@ -0,0 +1,124 @@ +import hashlib +import json +from pathlib import Path +import sys + +import pytest + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +import run_sfm_b1_r2_9arm as R + + +def _gpus(n=4): + return [ + R.GPU( + index=str(index), uuid=f"GPU-{index}", name="H100", memory_total_mib=95830, + memory_used_mib=20, utilization_percent=0, pci_bus_id=f"0000:{index:02x}:00.0", + ) + for index in range(n) + ] + + +def test_four_gpu_assignment_is_declared_three_two_two_two(): + allocation = R.assign_arms(list(R.arm_grid()), _gpus()) + by_index = { + gpu.index: [arm.name for arm in allocation[gpu.uuid]] for gpu in _gpus() + } + assert [len(by_index[str(index)]) for index in range(4)] == [3, 2, 2, 2] + assert {arm.replay_epochs for arm in allocation["GPU-0"]} == {1} + for index in ("1", "2", "3"): + assert { + arm.replay_epochs for arm in allocation[f"GPU-{index}"] + } == {10, 100} + assert sorted(sum(by_index.values(), [])) == sorted(arm.name for arm in R.arm_grid()) + + +def test_assignment_fails_when_capacity_is_insufficient(): + with pytest.raises(RuntimeError, match="exceed"): + R.assign_arms(list(R.arm_grid()), _gpus(2), max_arms_per_gpu=3) + + +def test_explicit_busy_gpu_fails_closed(): + gpus = _gpus(2) + with pytest.raises(RuntimeError, match="not idle"): + R.select_idle_gpus( + gpus, [{"gpu_uuid": "GPU-1"}], "0,1", + max_memory_mib=1024, max_utilization=5, + ) + selected = R.select_idle_gpus( + gpus, [{"gpu_uuid": "GPU-1"}], "auto", + max_memory_mib=1024, max_utilization=5, + ) + assert [gpu.index for gpu in selected] == ["0"] + + +def test_complete_marker_requires_contract_and_checkpoint_hashes(tmp_path): + arm = R.Arm(0.01, 10) + arm_dir = tmp_path / arm.name + arm_dir.mkdir() + history = [] + for round_i in range(3): + path = arm_dir / f"round_{round_i:02d}.pt" + path.write_bytes(f"round-{round_i}".encode()) + digest = hashlib.sha256(path.read_bytes()).hexdigest() + (arm_dir / f"round_{round_i:02d}.pt.COMPLETE.json").write_text(json.dumps( + dict(status="COMPLETE", path=str(path), sha256=digest) + )) + if round_i: + history.append(dict(round=round_i, checkpoint_sha256=digest)) + contract = R._expected_arm_contract( + arm, checkpoint_sha256="source", scene_profile="double_density_velocity_ood", + ell=R.ELL, cap=R.CAP, seed=7, verifier_workers=8, + ) + marker = dict( + status=R.ARM_STATUS, experiment=arm.name, + recipe=dict( + alpha=contract["alpha"], + replay_epochs=contract["replay_epochs"], + rounds=contract["rounds"], scene_profile=contract["scene_profile"], + seed=contract["seed"], verifier_workers=contract["verifier_workers"], + lr=contract["lr"], + ), + constants=dict(ell=contract["ell"], cap=contract["cap"]), + source_checkpoint_sha256="source", history=history, + ) + (arm_dir / "COMPLETE.json").write_text(json.dumps(marker)) + result = R.validate_complete_arm( + arm_dir, arm, checkpoint_sha256="source", + scene_profile="double_density_velocity_ood", + ell=R.ELL, cap=R.CAP, seed=7, verifier_workers=8, + ) + assert [row["round"] for row in result["checkpoints"]] == [0, 1, 2] + (arm_dir / "round_01.pt").write_bytes(b"changed") + with pytest.raises(RuntimeError, match="mismatch"): + R.validate_complete_arm( + arm_dir, arm, checkpoint_sha256="source", + scene_profile="double_density_velocity_ood", + ell=R.ELL, cap=R.CAP, seed=7, verifier_workers=8, + ) + + +def test_incomplete_nonempty_arm_is_not_overwritten(tmp_path): + arm = R.Arm(0.0, 1) + arm_dir = tmp_path / arm.name + arm_dir.mkdir() + (arm_dir / "partial.log").write_text("preserve") + with pytest.raises(RuntimeError, match="incomplete nonempty"): + R.validate_complete_arm( + arm_dir, arm, checkpoint_sha256="source", + scene_profile="double_density_velocity_ood", + ell=R.ELL, cap=R.CAP, seed=7, verifier_workers=8, + ) + + +def test_trainer_command_uses_complete_replay_epoch_cli(tmp_path): + checkpoint = tmp_path / "checkpoint.pt" + checkpoint.write_bytes(b"checkpoint") + args = type("Args", (), dict( + checkpoint=str(checkpoint), verifier_workers=8, seed=17, + ))() + command = R._trainer_command(args, R.Arm(0.1, 100), tmp_path / "arm") + assert command[command.index("--alpha") + 1] == "0.1" + assert command[command.index("--replay-epochs") + 1] == "100" + assert command[command.index("--verifier-workers") + 1] == "8" + assert "--adam-steps" not in command diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_b1_branch_compare_viz.py b/overnight_run_07_12_sfm/analysis/test_sfm_b1_branch_compare_viz.py new file mode 100644 index 0000000..99aa880 --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_sfm_b1_branch_compare_viz.py @@ -0,0 +1,52 @@ +import sfm_b1_branch_compare_viz as V +import sfm_b1_d_branch_viz as D +import sfm_b1_selector_compare_viz as S +import numpy as np + + +def test_summary_separates_window_labels_from_episode_outcomes(): + bundle = { + "traces": [ + {"executed_label": "verifier_positive"}, + {"executed_label": "verifier_negative"}, + {"executed_label": "verifier_positive"}, + ], + "outcomes": [ + {"success": True, "collision": False, "timeout": False}, + {"success": False, "collision": True, "timeout": False}, + ], + } + report = V.summarize(bundle) + assert report["contexts"] == 3 + assert report["executed_positive"] == 2 + assert report["executed_negative"] == 1 + assert report["executed_positive_fraction"] == 2 / 3 + assert (report["success"], report["collision"], report["timeout"]) == ( + 1, 1, 0, + ) + + +def test_selector_pair_requires_same_checkpoint_and_bank(): + common = { + "scenarios": [1, 2, 3], + "gammas": [.1, .2, .3, .4, .5, .7, 1.], + "environment": {"name": "test"}, + "sample_seed": 8, + "audit_seed": 9, + "checkpoint_sha256": "a" * 64, + } + S._validate_bundles((("margin", common), ("cost", dict(common)))) + different = dict(common, checkpoint_sha256="b" * 64) + try: + S._validate_bundles((("margin", common), ("cost", different))) + except ValueError as error: + assert "one pretrained checkpoint" in str(error) + else: + raise AssertionError("checkpoint mismatch must be rejected") + + +def test_robot_frame_uses_velocity_direction(): + trace = {"state": np.array([2., 3., 0., 2.])} + path = np.array([[2., 3.], [2., 4.], [3., 4.]]) + local = D._robot_frame(path, trace) + np.testing.assert_allclose(local, [[0., 0.], [1., 0.], [1., -1.]]) diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_b1_cost.py b/overnight_run_07_12_sfm/analysis/test_sfm_b1_cost.py index 9b9d1d3..59e0524 100644 --- a/overnight_run_07_12_sfm/analysis/test_sfm_b1_cost.py +++ b/overnight_run_07_12_sfm/analysis/test_sfm_b1_cost.py @@ -35,3 +35,35 @@ def test_gate_precedes_selector_and_nvp_when_none(): rows, selector="margin", state=np.zeros(4), ped_xy=np.array([[3., 3.]]), ped_vel=np.zeros((1, 2)), gamma=.5, ) is None + + +def test_balanced_rank_uses_rank_sum_with_safety_first_tie(monkeypatch): + rows = [] + for candidate_id in range(3): + controls = np.zeros((10, 2), np.float32) + controls[0, 0] = candidate_id + 1 + rows.append({ + "candidate_id": candidate_id, + "controls": controls, + "result": {"resolved": True, "y": 1}, + }) + monkeypatch.setattr( + C, "nominal_hp_margin", + lambda state, action, ped_xy, gamma: ( + 4.0 - float(action[0]), 1.0, 1.0, + ), + ) + monkeypatch.setattr( + C, "safemppi_proposal_cost", + lambda state, controls, goal, ped_xy, ped_vel: torch.tensor( + [4.0 - float(value) for value in controls[:, 0, 0]] + ), + ) + chosen = C.select_admissible( + rows, selector="balanced_rank", state=np.zeros(4), + ped_xy=np.array([[3., 3.]]), ped_vel=np.zeros((1, 2)), gamma=.5, + ) + assert chosen["candidate_id"] == 0 + assert chosen["safety_rank"] == 1 + assert chosen["performance_rank"] == 3 + assert chosen["rank_sum"] == 4 diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_b1_full_episode_audit.py b/overnight_run_07_12_sfm/analysis/test_sfm_b1_full_episode_audit.py new file mode 100644 index 0000000..2e481a2 --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_sfm_b1_full_episode_audit.py @@ -0,0 +1,76 @@ +import matplotlib.pyplot as plt +import numpy as np +import pytest + +import sfm_b1_full_episode_audit as A +import sfm_b1_full_episode_viz as V +import sfm_b1_viz as BV + + +def test_result_label_requires_a_resolved_full_h_positive(): + assert A._result_label(dict(resolved=False)) == "verifier_error" + assert A._result_label(dict(resolved=True, y=0)) == "verifier_negative" + assert A._result_label( + dict(resolved=True, y=1, full_h=True, terminal_step=10) + ) == "verifier_positive" + with pytest.raises(RuntimeError, match="full-H=10"): + A._result_label( + dict(resolved=True, y=1, full_h=False, terminal_step=3) + ) + + +def test_trap_is_a_separate_exact_ten_transition_event(): + short = [np.array([0.0, 0.0, 0.0, 0.0])] * 10 + assert not A._trap(short) + stationary = [np.array([0.0, 0.0, 0.0, 0.0])] * 11 + assert A._trap(stationary) + moving = [ + np.array([0.03 * step, 0.0, 0.0, 0.0]) for step in range(11) + ] + assert not A._trap(moving) + + +def test_nvp_does_not_remove_later_steps_from_the_trace_index(): + traces = [ + dict(scenario_id=7, gamma=.5, step=0, nvp_context=True), + dict(scenario_id=7, gamma=.5, step=1, nvp_context=False), + ] + index = V._index(traces) + assert sorted(index[(7, .5)]) == [0, 1] + with pytest.raises(ValueError, match="duplicate trace key"): + V._index(traces + [dict(traces[1])]) + + +def test_executed_color_is_the_verifier_label_not_the_nvp_event(): + assert V._executed_color( + dict(executed_label="verifier_positive", nvp_context=True) + ) == BV.BLUE + assert V._executed_color( + dict(executed_label="verifier_negative", nvp_context=False) + ) == BV.RED + + +def test_verifier_geometry_is_drawn_only_for_positive_executed_window(monkeypatch): + calls = [] + monkeypatch.setattr( + V.DV, "checked_verifier_levels", + lambda trace, query, H: calls.append(H) or dict( + polygons=[], outer_polygon=None + ), + ) + monkeypatch.setattr(V.DV, "_draw_verifier_geometry", lambda axis, audit: None) + base = dict( + state=np.zeros(4), next_state=np.zeros(4), + executed_controls=np.zeros((10, 2)), + executed_result=dict(segment=np.zeros((11, 2))), + ) + figure, axis = plt.subplots() + V._draw_executed( + axis, dict(base, executed_label="verifier_negative") + ) + assert calls == [] + V._draw_executed( + axis, dict(base, executed_label="verifier_positive") + ) + assert calls == [10] + plt.close(figure) diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_b1_kazuki_repair.py b/overnight_run_07_12_sfm/analysis/test_sfm_b1_kazuki_repair.py new file mode 100644 index 0000000..8670e05 --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_sfm_b1_kazuki_repair.py @@ -0,0 +1,283 @@ +from types import SimpleNamespace + +import numpy as np +import torch + +import sfm_b1_kazuki_repair as R +import sfm_b1_kazuki_repair_audit as A +import sfm_scene as SS + + +class _Policy: + d = 20 + H_pred = 10 + u_max = 2.0 + + +class _DeterministicGuidancePolicy: + H_pred = 10 + d = 20 + u_max = 2.0 + + def _expand_ctx(self, context, count): + return context.reshape(1, -1).expand(int(count), -1) + + def forward(self, value, tau, context): + del tau, context + return 0.15 * torch.tanh(value) + + +def test_locked_guidance_config_is_exact_and_b_sized(): + config = R.locked_guidance_config(0.1, n_sample=4) + assert config.safe_coefs == (0.3,) + assert config.goal_coef == 0.5 + assert config.safe_coef_gamma_span == 0.0 + assert config.goal_coef_gamma_span == 0.0 + assert config.n_sample == 4 + assert R.collector_ode_times(8) == tuple(index / 8 for index in range(9)) + + +def test_same_latent_repair_reuses_exact_x0_and_has_no_refinement(monkeypatch): + seen = {} + + def fake_predict(ped_xy, ped_vel, H, dt, device, dtype): + del ped_xy, ped_vel, dt + return torch.zeros(1, H, 1, 2, device=device, dtype=dtype) + + def fake_guided( + policy, context, state, goal, ped_prediction, ped_velocity, + radius, x0, times, config, collect_diagnostics, + ): + del ( + policy, context, state, goal, ped_prediction, ped_velocity, + radius, + ) + seen.update( + x0=x0.detach().clone(), + times=tuple(times), + n_sample=config.n_sample, + collect=collect_diagnostics, + ) + trace = [dict( + integrated_goal_guidance=torch.ones_like(x0), + integrated_safety_guidance=2 * torch.ones_like(x0), + component_semantics="test", + )] + return x0 + 0.1, trace, x0 + + monkeypatch.setattr(R.KZ, "predict_pedestrians_t", fake_predict) + monkeypatch.setattr(R.KZ, "guided_generate", fake_guided) + x0 = torch.arange(80, dtype=torch.float32).reshape(4, 20) / 100 + context = torch.zeros(37) + controls, diagnostics = R.same_latent_guided_controls( + _Policy(), + context, + np.zeros(4, np.float32), + np.zeros((1, 2), np.float32), + np.zeros((1, 2), np.float32), + 0.5, + x0, + nfe=8, + ) + torch.testing.assert_close(seen["x0"], x0) + assert seen["times"] == tuple(index / 8 for index in range(9)) + assert seen["n_sample"] == 4 and seen["collect"] + assert tuple(controls.shape) == (4, 10, 2) + assert diagnostics["no_new_latents"] + assert diagnostics["no_mppi_refinement"] + np.testing.assert_allclose( + diagnostics["goal_first_action"], np.full((4, 2), 2.0), + ) + np.testing.assert_allclose( + diagnostics["safety_first_action"], np.full((4, 2), 4.0), + ) + + +def test_guidance_diagnostics_do_not_change_guided_controls(): + torch.manual_seed(31) + policy = _DeterministicGuidancePolicy() + context = torch.zeros(3) + state = np.zeros(4, np.float32) + goal = torch.tensor(SS.GOAL, dtype=torch.float32) + x0 = torch.randn(4, 20) + ped_prediction = torch.full((1, 10, 2), 20.0) + ped_velocity = torch.zeros(1, 2) + config = R.locked_guidance_config(0.5, n_sample=4) + times = R.collector_ode_times(8) + plain, _, _ = R.KZ.guided_generate( + policy, + context, + state, + goal, + ped_prediction, + ped_velocity, + SS.R_PED + config.collision_margin, + x0.clone(), + times, + config, + collect_diagnostics=False, + ) + diagnostic, trace, _ = R.KZ.guided_generate( + policy, + context, + state, + goal, + ped_prediction, + ped_velocity, + SS.R_PED + config.collision_margin, + x0.clone(), + times, + config, + collect_diagnostics=True, + ) + torch.testing.assert_close(diagnostic, plain, rtol=0, atol=0) + assert "integrated_goal_guidance" in trace[-1] + assert "integrated_safety_guidance" in trace[-1] + + +def test_predictive_trap_matches_post_action_predicate_and_stops_at_three(): + states = [ + np.array([0.01 * index, 0.0, 0.0, 0.0], np.float32) + for index in range(10) + ] + action = np.zeros(2, np.float32) + assert R.predicted_trap(states, action) + assert R.post_action_trap_matches(states, action) + streak = 0 + for expected in (1, 2): + streak, stop = R.next_trap_streak(streak, True) + assert streak == expected and not stop + streak, stop = R.next_trap_streak(streak, True) + assert streak == 3 and stop + assert R.next_trap_streak(streak, False) == (0, False) + + +def test_trigger_prefers_existing_trap_then_nvp_then_prediction(monkeypatch): + replica = SimpleNamespace(states=[np.zeros(4, np.float32)]) + monkeypatch.setattr(A.FA, "_trap", lambda states: True) + assert A._repair_trigger(replica, {"controls": np.zeros((10, 2))}) == "trap_streak" + monkeypatch.setattr(A.FA, "_trap", lambda states: False) + assert A._repair_trigger(replica, None) == "finite_B_NVP" + monkeypatch.setattr(A.KR, "predicted_trap", lambda states, action: True) + assert A._repair_trigger( + replica, {"controls": np.zeros((10, 2))}, + ) == "predicted_trap" + monkeypatch.setattr(A.KR, "predicted_trap", lambda states, action: False) + assert A._repair_trigger( + replica, {"controls": np.zeros((10, 2))}, + ) is None + + +def test_manifest_excludes_privileged_and_independent_fallbacks(): + manifest = R.manifest() + assert manifest["safe_coef"] == 0.3 + assert manifest["goal_coef"] == 0.5 + assert "no privileged MPC" in manifest["exclusions"] + assert "no independent raw fallback" in manifest["exclusions"] + assert SS.DT > 0 + + +def _negative_rows(): + rows = [] + for candidate_id in range(4): + rows.append(dict( + candidate_id=16 + candidate_id, + parent_candidate_id=candidate_id, + acquisition_step=candidate_id, + controls=np.full((10, 2), candidate_id, np.float32), + x0=np.full(20, candidate_id, np.float32), + sigma=0.1 + candidate_id, + mode="test", + query_id=candidate_id, + result=dict( + resolved=True, + y=0, + full_h=True, + terminal_step=10, + taskspace=True, + collision_free=True, + certificate=False, + diagnostics={}, + ), + )) + return rows + + +def _prepared(): + return dict( + state=np.zeros(4, np.float32), + hp10=torch.zeros(10, 16, 12), + low=torch.zeros(5), + hist=torch.zeros(16, 2), + ped_xy=np.zeros((1, 2), np.float32), + ped_vel=np.zeros((1, 2), np.float32), + ) + + +def test_neutral_margin_ranks_all_four_negatives_without_relabeling( + monkeypatch, +): + rows = _negative_rows() + margins = iter((0.1, 0.7, 0.7, -0.2)) + monkeypatch.setattr( + A.BC, + "nominal_hp_margin", + lambda *args: (next(margins), 1.0, 1.0), + ) + chosen = A._select_neutral(rows, "margin", _prepared(), 0.5) + assert chosen["candidate_id"] == 17 + assert all(row["result"]["y"] == 0 for row in rows) + + +def test_neutral_cost_ranks_all_four_negatives(monkeypatch): + rows = _negative_rows() + monkeypatch.setattr( + A.BC, + "nominal_hp_margin", + lambda *args: (0.2, 1.0, 1.0), + ) + monkeypatch.setattr( + A.BC, + "safemppi_proposal_cost", + lambda *args, **kwargs: torch.tensor([3.0, 1.0, 1.0, 2.0]), + ) + chosen = A._select_neutral( + rows, "safemppi_cost", _prepared(), 0.5, + ) + assert chosen["candidate_id"] == 17 + assert chosen["expert_cost"] == 1.0 + + +def test_neutral_selection_requires_four_exact_full_h_negatives(): + rows = _negative_rows() + rows[0]["result"]["resolved"] = False + assert A._select_neutral(rows, "margin", _prepared(), 0.5) is None + rows = _negative_rows() + rows[0]["result"]["y"] = 1 + assert A._select_neutral(rows, "margin", _prepared(), 0.5) is None + + +def test_neutral_record_is_separate_and_nontraining(tmp_path): + rows = _negative_rows() + for row in rows: + row["hp_margin"] = 0.2 + replica = SimpleNamespace(scenario_id=250001, gamma=0.5) + record = A._neutral_record( + 0, + replica, + _prepared(), + rows[0], + step=9, + selector="margin", + repair_trigger="finite_B_NVP", + ) + assert record["semantic_label"] == "neutral" + assert record["verifier_y"] == 0 + assert not record["train_eligible"] + assert not record["replay_default"] + assert not record["gp_eligible"] + path = tmp_path / "neutral_round.pt" + marker = A._save_neutral_records(path, [record]) + payload = torch.load(path, map_location="cpu", weights_only=False) + assert marker["D0"] == 1 + assert payload["records"][0]["population"] == "D0" diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_b1_neutral_multiround.py b/overnight_run_07_12_sfm/analysis/test_sfm_b1_neutral_multiround.py new file mode 100644 index 0000000..6e66052 --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_sfm_b1_neutral_multiround.py @@ -0,0 +1,402 @@ +from types import SimpleNamespace + +import json +import numpy as np +import pytest +import torch + +import sfm_b1_kazuki_repair_audit as RA +import sfm_b1_neutral_multiround as M +import sfm_b1_offline_store as OS +import sfm_protocol as SP + + +class _TinyPolicy(torch.nn.Module): + def __init__(self): + super().__init__() + self.enc_grid = torch.nn.Linear(1, 1, bias=False) + self.enc_grid.requires_grad_(False) + self.scale = torch.nn.Parameter(torch.tensor(0.1)) + + def ctx_from(self, grid, low, hist): + return low[:, :1] + + def cfm_loss(self, controls, context, weights=None): + target = controls.reshape(len(controls), -1).mean(dim=1) + prediction = self.scale * context[:, 0] + per = (prediction - target).square() + return per.mean() if weights is None else (per * weights).mean() + + def phi_s_from_x0(self, controls, context, x0, s): + return x0 + + +def _records(population="D0"): + holder = M._NeutralHolder(1) + rows = [] + for index, gamma in enumerate((0.1, 1.0)): + holder.contexts.append({ + "context_id": index, + "round": 1, + "scenario_id": 260000 + index, + "gamma": gamma, + "step": index, + "hp10": np.zeros((10, 16, 12), np.float32), + "low5": np.asarray([1 + index, 0, 0, 0, gamma], np.float32), + "hist": np.zeros((16, 2), np.float32), + }) + row = { + "query_id": index, + "window_id": index, + "context_id": index, + "controls": np.full((10, 2), 0.2 + index, np.float32), + "x0": np.zeros(20, np.float32), + "y": int(population == "Dplus"), + } + holder.windows.append(row) + rows.append((holder, row)) + return rows + + +def test_study_config_pins_two_scenarios_and_protocol(): + assert M.StudyConfig(name="x").validate().scenarios_per_round == 2 + with pytest.raises(ValueError, match="two scenarios"): + M.StudyConfig(name="x", scenarios_per_round=8).validate() + with pytest.raises(ValueError, match="K/B/H/T"): + M.StudyConfig(name="x", K=64).validate() + + +def test_resume_accepts_only_semantically_identical_legacy_defaults(): + current = M.StudyConfig(name="continued", rounds=100).__dict__ + previous = dict(current) + previous["name"] = "legacy" + previous["rounds"] = 50 + previous.pop("encoder_lr_ratio") + previous.pop("neutral_replay") + normalized = M._validate_resume_config(previous, current) + assert normalized == {"encoder_lr_ratio": 0.0, "neutral_replay": True} + + current_no_d0 = dict(current, neutral_replay=False) + with pytest.raises(RuntimeError, match="neutral_replay"): + M._validate_resume_config(previous, current_no_d0) + + +def test_population_update_uses_every_row_once_per_inner_pass(): + records = _records() + policy = _TinyPolicy() + optimizer = torch.optim.Adam([policy.scale], lr=3.0e-5) + before = float(policy.scale.detach()) + report = M._population_update( + policy, + optimizer, + records, + population="D0", + inner_steps=4, + batch=1, + device="cpu", + seed=17, + ) + assert report["optimizer_steps"] == 4 + assert report["sample_exposures"] == 8 + assert report["exact_once_per_inner_step"] + assert len(report["exposure_identity_sha256"]) == 4 + assert float(policy.scale.detach()) != before + + +def test_restore_optimizer_preserves_global_adam_step(tmp_path): + policy = _TinyPolicy() + optimizer = torch.optim.Adam([policy.scale], lr=3.0e-5) + for _ in range(4): + optimizer.zero_grad(set_to_none=True) + policy.scale.square().backward() + optimizer.step() + path = tmp_path / "optimizer.pt" + torch.save({"round": 2, "optimizer": optimizer.state_dict()}, path) + + resumed_policy = _TinyPolicy() + resumed = torch.optim.Adam([resumed_policy.scale], lr=3.0e-5) + report = M._restore_optimizer( + resumed, + str(path), + [resumed_policy.scale], + resume_round=2, + inner_steps=1, + ) + assert report["expected_adam_step"] == 4 + assert int(resumed.state[resumed_policy.scale]["step"].item()) == 4 + + with pytest.raises(RuntimeError, match="round mismatch"): + M._restore_optimizer( + torch.optim.Adam([resumed_policy.scale], lr=3.0e-5), + str(path), + [resumed_policy.scale], + resume_round=3, + inner_steps=1, + ) + + one_phase = torch.optim.Adam([resumed_policy.scale], lr=3.0e-5) + one_phase_report = M._restore_optimizer( + one_phase, + str(path), + [resumed_policy.scale], + resume_round=2, + inner_steps=2, + neutral_replay=False, + ) + assert one_phase_report["expected_adam_step"] == 4 + + +def test_nested_resume_round_refs_are_recursive_and_contiguous(tmp_path): + deliveries = [] + lineage = [] + for round_i in (1, 2, 3): + checkpoint = tmp_path / f"round_{round_i:02d}.pt" + post_positive = tmp_path / f"round_{round_i:02d}_post_positive.pt" + checkpoint.write_bytes(f"checkpoint-{round_i}".encode()) + post_positive.write_bytes(f"positive-{round_i}".encode()) + marker = tmp_path / f"round_{round_i:02d}.json" + record = { + "status": M.ROUND_STATUS, "round": round_i, + "checkpoint": str(checkpoint), + "checkpoint_sha256": M.FA._sha256_file(checkpoint), + "scenarios": [260000 + 2 * round_i - 2, 260000 + 2 * round_i - 1], + } + marker.write_text(json.dumps(record)) + current_ref = M._round_record_ref(marker, record) + delivery = { + "status": M.STATUS, + "round_records": [str(marker)], + "round_record_refs": [current_ref], + } + if deliveries: + prior_path, prior = deliveries[-1] + delivery["resume"] = { + "delivery": str(prior_path), + "delivery_sha256": M.FA._sha256_file(prior_path), + "round_record_refs": list(lineage), + } + delivery_path = tmp_path / f"delivery_{round_i}.json" + delivery_path.write_text(json.dumps(delivery)) + lineage.append(current_ref) + deliveries.append((delivery_path, delivery)) + + final_path, final_delivery = deliveries[-1] + refs = M._delivery_lineage_refs(final_delivery, final_path) + assert [row["round"] for row in refs] == [1, 2, 3] + + broken = dict(deliveries[-1][1]) + broken["round_record_refs"] = [dict(lineage[-1], path=str(tmp_path / "wrong"))] + with pytest.raises(RuntimeError, match="current resume round snapshot"): + M._delivery_lineage_refs(broken, deliveries[-1][0]) + + +def _positive_result(): + return { + "resolved": True, + "y": 1, + "taskspace": True, + "collision_free": True, + "certificate": True, + "full_h": True, + "terminal_step": 10, + "train_eligible": True, + "diagnostics": {}, + } + + +def _previous_shard(counts): + shard = OS.ExecutedRoundShard(1) + for gamma, count in zip(SP.GAMMAS, counts): + for index in range(int(count)): + context_id = shard.add_context( + scenario_id=100000 + index, + gamma=gamma, + step=index, + state=np.zeros(4, np.float32), + hp10=np.zeros((10, 16, 12), np.float32), + low5=np.zeros(5, np.float32), + hist=np.zeros((16, 2), np.float32), + ped_xy=np.zeros((1, 2), np.float32), + ped_vel=np.zeros((1, 2), np.float32), + ) + shard.add_executed_window( + context_id, + np.zeros((10, 2), np.float32), + np.full(20, index + float(gamma), np.float32), + _positive_result(), + execution_source="test", + nvp_context=False, + ) + return shard + + +def test_dynamic_gp_uses_equal_supported_quota_without_backfill(): + previous = _previous_shard((3, 5, 5, 5, 5, 5, 5)) + gp, identities, selection = RA._gamma_balanced_gp( + _TinyPolicy(), + previous, + gammas=tuple(map(float, SP.GAMMAS)), + round_i=2, + ell=RA.DEFAULT_ELL, + cap=512, + lam=1.0e-2, + phi_s=0.9, + device="cpu", + seed=5, + ) + assert selection["quota"] == 3 + assert selection["effective_cap"] == 22 + assert selection["per_gamma"]["0.1"] == 3 + assert sum(selection["per_gamma"].values()) == 22 + assert len(identities) == len(set(identities)) == 22 + assert gp.diagnostics()["n"] == 22 + + +def test_dynamic_gp_is_empty_on_round_one(): + gp, identities, selection = RA._gamma_balanced_gp( + _TinyPolicy(), + None, + gammas=tuple(map(float, SP.GAMMAS)), + round_i=1, + ell=RA.DEFAULT_ELL, + cap=512, + lam=1.0e-2, + phi_s=0.9, + device="cpu", + seed=5, + ) + assert identities == [] + assert selection["effective_cap"] == 0 + assert gp.diagnostics()["n"] == 0 + + +def test_dynamic_gp_fails_if_any_declared_gamma_has_no_support(): + previous = _previous_shard((0, 2, 2, 2, 2, 2, 2)) + with pytest.raises(RuntimeError, match="cannot support all declared"): + RA._gamma_balanced_gp( + _TinyPolicy(), + previous, + gammas=tuple(map(float, SP.GAMMAS)), + round_i=2, + ell=RA.DEFAULT_ELL, + cap=512, + lam=1.0e-2, + phi_s=0.9, + device="cpu", + seed=5, + ) + + +def test_eval_round_parser_is_staged_and_requires_final(): + assert M._parse_eval_rounds( + "0,1,2,5,10,20,30,40,50", 50, + ) == (0, 1, 2, 5, 10, 20, 30, 40, 50) + with pytest.raises(ValueError, match="final round"): + M._parse_eval_rounds("0,1,2", 50) + + +def _probe_row( + anchor_id, + *, + population, + nvp, + progress, + admissible, + rmse, +): + return { + "anchor_id": anchor_id, + "gamma": 0.5, + "trigger": "finite_B_NVP", + "population": population, + "B_orig_NVP": nvp, + "target_full_rmse": rmse, + "selected_one_step_progress": progress, + "all_K_admissible": admissible, + "B_orig_admissible": int(not nvp), + "target_latent_admissible": bool(not nvp), + "_feature": np.asarray([1.0, 0.0], np.float32), + "_windows": np.zeros((16, 10, 2), np.float32), + } + + +def test_probe_classifies_repair_stall_imitation_and_regression(): + before = [ + _probe_row( + 0, population="D0", nvp=True, progress=None, + admissible=0, rmse=1.0, + ), + _probe_row( + 1, population="D0", nvp=True, progress=None, + admissible=0, rmse=1.0, + ), + _probe_row( + 2, population="D0", nvp=True, progress=None, + admissible=0, rmse=1.0, + ), + _probe_row( + 3, population="Dplus", nvp=False, progress=0.1, + admissible=3, rmse=0.0, + ), + ] + after = [ + _probe_row( + 0, population="D0", nvp=False, progress=0.2, + admissible=2, rmse=0.4, + ), + _probe_row( + 1, population="D0", nvp=False, progress=-0.1, + admissible=1, rmse=0.4, + ), + _probe_row( + 2, population="D0", nvp=True, progress=None, + admissible=0, rmse=0.4, + ), + _probe_row( + 3, population="Dplus", nvp=True, progress=None, + admissible=0, rmse=0.2, + ), + ] + summary = M._probe_comparison(before, after)["pooled"] + assert summary["D0_actual_repairs"] == 1 + assert summary["D0_safe_stalls"] == 1 + assert summary["D0_imitation_only"] == 1 + assert summary["Dplus_regressed"] == 1 + + +def test_neutral_record_accepts_explicit_round(): + result = { + "resolved": True, + "y": 0, + "full_h": True, + "terminal_step": 10, + } + chosen = { + "controls": np.zeros((10, 2), np.float32), + "x0": np.zeros(20, np.float32), + "result": result, + "candidate_id": 16, + "sigma": 0.1, + "hp_margin": 0.2, + } + prepared = { + "state": np.zeros(4, np.float32), + "hp10": torch.zeros(10, 16, 12), + "low": torch.zeros(5), + "hist": torch.zeros(16, 2), + "ped_xy": np.zeros((1, 2), np.float32), + "ped_vel": np.zeros((1, 2), np.float32), + } + record = RA._neutral_record( + 0, + SimpleNamespace(scenario_id=1, gamma=0.5), + prepared, + chosen, + step=3, + round_i=2, + selector="margin", + repair_trigger="finite_B_NVP", + ) + assert record["round"] == 2 + assert record["verifier_y"] == 0 + assert not record["gp_eligible"] diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_b1_neutral_teacher_sanity.py b/overnight_run_07_12_sfm/analysis/test_sfm_b1_neutral_teacher_sanity.py new file mode 100644 index 0000000..c18cb58 --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_sfm_b1_neutral_teacher_sanity.py @@ -0,0 +1,168 @@ +import copy + +import numpy as np +import pytest +import torch + +import sfm_b1_neutral_teacher_sanity as N +import sfm_b1_store as BS + + +def _payload(): + records = [] + for index, gamma in enumerate((0.1, 1.0)): + records.append({ + "neutral_id": index, + "population": "D0", + "semantic_label": "neutral", + "round": 1, + "scenario_id": 250_000 + index, + "gamma": gamma, + "step": index, + "state": np.zeros(4, np.float32), + "hp10": np.zeros((10, 16, 12), np.float32), + "low5": np.asarray([1 + index, 0, 0, 0, gamma], np.float32), + "hist": np.zeros((16, 2), np.float32), + "ped_xy": np.zeros((40, 2), np.float32), + "ped_vel": np.zeros((40, 2), np.float32), + "controls": np.full((10, 2), 0.2 + index, np.float32), + "x0": np.zeros(20, np.float32), + "verifier_result": { + "resolved": True, + "y": 0, + "full_h": True, + "terminal_step": 10, + }, + "verifier_y": 0, + "train_eligible": False, + "replay_default": False, + "gp_eligible": False, + }) + return { + "status": "SFM_B1_NEUTRAL_ROUND_COMPLETE", + "round": 1, + "records": records, + "summary": {"D0": len(records)}, + } + + +class _TinyPolicy(torch.nn.Module): + def __init__(self): + super().__init__() + self.enc_grid = torch.nn.Linear(1, 1, bias=False) + self.enc_grid.requires_grad_(False) + self.scale = torch.nn.Parameter(torch.tensor(0.1)) + + def ctx_from(self, grid, low, hist): + return low[:, :1] + + def cfm_loss(self, controls, context, weights=None): + target = controls.reshape(len(controls), -1).mean(dim=1) + prediction = self.scale * context[:, 0] + per = (prediction - target).square() + return per.mean() if weights is None else (per * weights).mean() + + +def test_neutral_conversion_preserves_exact_negative_isolation(): + payload = _payload() + original = copy.deepcopy(payload) + holder, records = N._neutral_records(payload) + assert len(holder.contexts) == len(records) == 2 + assert all(row["y"] == 0 for _, row in records) + assert all(row["semantic_label"] == "neutral" for _, row in records) + mass, accounting = BS.hierarchy_mass(records) + assert sum(mass.values()) == pytest.approx(1.0) + assert accounting["gamma"] == pytest.approx({"0.1": 0.5, "1.0": 0.5}) + for source in payload["records"]: + assert not source["train_eligible"] + assert not source["replay_default"] + assert not source["gp_eligible"] + assert source["verifier_y"] == 0 + for actual, expected in zip(payload["records"], original["records"]): + assert actual.keys() == expected.keys() + np.testing.assert_array_equal(actual["controls"], expected["controls"]) + np.testing.assert_array_equal(actual["x0"], expected["x0"]) + + +def test_neutral_update_uses_whole_support_once_per_inner_step(): + _, records = N._neutral_records(_payload()) + policy = _TinyPolicy() + optimizer = torch.optim.Adam( + [policy.scale], lr=1.0e-2 + ) + before = float(policy.scale.detach()) + report = N._neutral_update( + policy, + optimizer, + records, + steps=4, + batch=1, + device="cpu", + seed=7, + ) + assert report["optimizer_steps"] == 4 + assert report["sample_exposures"] == 8 + assert report["exact_once_per_inner_step"] + assert report["stored_x0_used_for_training"] is False + assert len(report["losses"]) == 4 + assert float(policy.scale.detach()) != before + assert all(row["y"] == 0 for _, row in records) + + +def test_neutral_conversion_rejects_relabeling(): + payload = _payload() + payload["records"][0]["verifier_y"] = 1 + with pytest.raises(RuntimeError, match="D0 semantics"): + N._neutral_records(payload) + + +def test_chunked_weighting_gradient_matches_direct_hierarchy_objective(): + _, records = N._neutral_records(_payload()) + mass, _ = BS.hierarchy_mass(records) + policy = _TinyPolicy() + policy.zero_grad(set_to_none=True) + N._objective( + policy, + records, + mass, + batch=1, + device="cpu", + seed=11, + backward=True, + ) + chunked = policy.scale.grad.detach().clone() + + policy.zero_grad(set_to_none=True) + direct = 0.0 + for holder, row in records: + context = holder.contexts[row["context_id"]] + target = torch.as_tensor(row["controls"]).mean() + low = torch.as_tensor(context["low5"]) + direct = direct + mass[(id(holder), row["query_id"])] * ( + policy.scale * low[0] - target + ).square() + direct.backward() + torch.testing.assert_close(chunked, policy.scale.grad) + + +def test_gradient_temporarily_enables_training_and_restores_mode(): + _, records = N._neutral_records(_payload()) + policy = _TinyPolicy() + policy.eval() + gradient = N._gradient( + policy, + records, + batch=1, + device="cpu", + seed=13, + ) + assert gradient["scale"] is not None + assert not policy.training + assert policy.scale.grad is None + + +def test_neutral_conversion_rejects_nested_verifier_inconsistency(): + payload = _payload() + payload["records"][0]["verifier_result"]["full_h"] = False + with pytest.raises(RuntimeError, match="D0 semantics"): + N._neutral_records(payload) diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_b1_offline_18arm_compare.py b/overnight_run_07_12_sfm/analysis/test_sfm_b1_offline_18arm_compare.py new file mode 100644 index 0000000..9ead484 --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_sfm_b1_offline_18arm_compare.py @@ -0,0 +1,66 @@ +import json + +import sfm_b1_offline_18arm_compare as C + + +def _delivery(root, selector): + aggregate = root / "evaluation" / "aggregate" + aggregate.mkdir(parents=True) + rows = [] + prefix = { + "margin": "offline_exec", + "safemppi_cost": "offline_exec_safemppi_cost", + "balanced_rank": "offline_exec_balanced_rank", + }[selector] + for alpha in (0.0, 0.01, 0.1): + for exposure in (1, 10, 100): + arm = ( + f"{prefix}_alpha{str(alpha).replace('.', 'p')}_" + f"exposures{exposure:03d}" + ) + for round_i in range(11): + rows.append({ + "selector": selector, + "arm": arm, + "alpha": alpha, + "exposure_epochs": exposure, + "round": round_i, + "SR": .5, "CR": .5, "timeout": 0., + "Validity": .4, "clearance": .1, "time_to_goal": 9., + }) + (root / "DELIVERY_COMPLETE.json").write_text(json.dumps({ + "status": "SFM_B1_OFFLINE_9ARM_DELIVERY_COMPLETE", + "contract": {"execution_selector": selector}, + })) + (aggregate / "AGGREGATE_COMPLETE.json").write_text(json.dumps({ + "status": "SFM_B1_OFFLINE_9ARM_AGGREGATE_COMPLETE", + "rows": rows, + })) + + +def test_compare_requires_and_combines_paired_99_row_sweeps(tmp_path): + margin = tmp_path / "margin" + cost = tmp_path / "cost" + _delivery(margin, "margin") + _delivery(cost, "safemppi_cost") + result = C.compare(margin, cost, tmp_path / "comparison") + assert result["status"] == C.STATUS + assert result["rows"] == 198 + assert result["paired_r0"]["CR"] == .5 + assert (tmp_path / "comparison" / "paired_18arm_raw_m50.png").is_file() + + +def test_compare_accepts_third_balanced_selector_sweep(tmp_path): + margin = tmp_path / "margin" + cost = tmp_path / "cost" + balanced = tmp_path / "balanced" + _delivery(margin, "margin") + _delivery(cost, "safemppi_cost") + _delivery(balanced, "balanced_rank") + result = C.compare( + margin, cost, tmp_path / "comparison", + balanced_root=balanced, + ) + assert result["rows"] == 297 + assert result["balanced_rank_root"] == str(balanced.resolve()) + assert (tmp_path / "comparison" / "paired_27arm_raw_m50.png").is_file() diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_b1_offline_eval.py b/overnight_run_07_12_sfm/analysis/test_sfm_b1_offline_eval.py new file mode 100644 index 0000000..f584705 --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_sfm_b1_offline_eval.py @@ -0,0 +1,258 @@ +from __future__ import annotations + +import json +from pathlib import Path + +import numpy as np +import pytest +import torch + +import sfm_b1_offline_eval as E +import sfm_protocol as SP + + +def _trajectory(n_steps=3, *, collision=False): + controls = np.arange(n_steps * 2, dtype=np.float32).reshape(n_steps, 2) + return { + "episode": 1, + "gamma": 0.5, + "status": "collision" if collision else "success", + "success": not collision, + "collision": collision, + "timeout": False, + "steps": n_steps, + "time_to_goal": 1.0 if not collision else None, + "successful_clearance": 0.2 if not collision else None, + "states": np.zeros((n_steps + 1, 4), np.float32), + "controls": controls, + "ped_xy": np.zeros((n_steps, 0, 2), np.float32), + "ped_vel": np.zeros((n_steps, 0, 2), np.float32), + } + + +def _compact_row(episode, gamma, validity): + evaluated = 10 + valid = int(round(float(validity) * evaluated)) + return { + "episode": int(episode), + "gamma": float(gamma), + "status": "success", + "success": True, + "collision": False, + "timeout": False, + "time_to_goal": 9.0, + "successful_clearance": 0.2, + "validity": float(validity), + "valid_windows": valid, + "evaluated_windows": evaluated, + "verifier_errors": 0, + } + + +def test_terminal_windows_use_actual_executed_controls_and_all_starts(monkeypatch): + row = _trajectory(12) + calls = [] + + def fake(state, controls, ped_xy, ped_vel, gamma): + calls.append(np.asarray(controls).copy()) + return { + "resolved": True, + "y": 1, + "window_horizon": len(controls), + } + + monkeypatch.setattr(E.SM, "verify_executed_window", fake) + result = E._verify_executed_episode(row) + + assert [len(controls) for controls in calls] == [ + 10, 10, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1, + ] + assert np.array_equal(calls[0], row["controls"][:10]) + assert np.array_equal(calls[-1], row["controls"][-1:]) + assert result == { + "validity": 1.0, + "valid_windows": 12, + "evaluated_windows": 12, + "verifier_errors": 0, + } + + +def test_validity_is_fractional_and_does_not_stop_at_first_negative(monkeypatch): + row = _trajectory(3, collision=True) + outcomes = iter((1, 0, 1)) + + def fake(state, controls, ped_xy, ped_vel, gamma): + return { + "resolved": True, + "y": next(outcomes), + "window_horizon": len(controls), + } + + monkeypatch.setattr(E.SM, "verify_executed_window", fake) + result = E._verify_executed_episode(row) + assert result["validity"] == pytest.approx(2 / 3) + assert result["valid_windows"] == 2 + assert result["evaluated_windows"] == 3 + assert result["verifier_errors"] == 0 + + +def test_verifier_error_is_not_silently_counted_as_negative(monkeypatch): + row = _trajectory(3) + calls = 0 + + def fake(state, controls, ped_xy, ped_vel, gamma): + nonlocal calls + calls += 1 + if calls == 2: + return {"resolved": False, "error": "solver failed"} + return {"resolved": True, "y": 1, "window_horizon": len(controls)} + + monkeypatch.setattr(E.SM, "verify_executed_window", fake) + result = E._verify_executed_episode(row) + assert result["verifier_errors"] == 1 + assert result["evaluated_windows"] == 1 + + +def test_zero_transition_trajectory_has_defined_zero_validity(): + result = E._verify_executed_episode(_trajectory(0)) + assert result == { + "validity": 0.0, + "valid_windows": 0, + "evaluated_windows": 0, + "verifier_errors": 0, + } + + +def test_summary_uses_mean_of_per_trajectory_fractions(): + rows = [ + _compact_row(1, .5, 1.0), + _compact_row(2, .5, .5), + ] + summary = E._summarize_one(rows, seed=7) + assert summary["Validity"]["mean"] == pytest.approx(.75) + assert summary["Validity"]["valid_windows"] == 15 + assert summary["Validity"]["evaluated_windows"] == 20 + assert summary["Validity"]["window_weighted_fraction"] == pytest.approx(.75) + assert "V_safe" not in summary + + +def test_temperature_defaults_to_one_and_is_validated(): + parser = E.build_parser() + args = parser.parse_args([ + "--checkpoints", "a.pt", + "--labels", "r0", + "--output-dir", "out", + ]) + assert args.temperature == 1.0 + assert args.temperature_by_gamma is None + + args.temperature = 0.0 + args.m_per_gamma = 1 + with pytest.raises(ValueError, match="temperature"): + E.run(args) + + +def test_temperature_is_part_of_noise_and_cache_contract(monkeypatch, tmp_path): + monkeypatch.setattr(E, "M_PER_GAMMA", 2) + monkeypatch.setattr(E, "TEMPERATURE", 0.7) + monkeypatch.setattr( + E, "TEMPERATURE_BY_GAMMA", tuple(.7 for _ in SP.GAMMAS) + ) + _, metadata = E._noise_bank(ep0=10, d=4, seed=12) + assert metadata["temperature"] == pytest.approx(0.7) + assert metadata["temperature_by_gamma"] == pytest.approx( + [.7] * len(SP.GAMMAS) + ) + + checkpoint = tmp_path / "checkpoint.pt" + checkpoint.write_bytes(b"checkpoint") + monkeypatch.setattr(E, "_sha256_file", lambda path: "evaluator") + key_07 = E._cell_key( + checkpoint_sha256="checkpoint", + scene_profile="double_density_velocity_ood", + ep0=10, + noise_meta=metadata, + ) + monkeypatch.setattr(E, "TEMPERATURE", 1.0) + monkeypatch.setattr( + E, "TEMPERATURE_BY_GAMMA", tuple(1.0 for _ in SP.GAMMAS) + ) + key_10 = E._cell_key( + checkpoint_sha256="checkpoint", + scene_profile="double_density_velocity_ood", + ep0=10, + noise_meta={**metadata, "temperature": 1.0}, + ) + assert key_07 != key_10 + + +def test_gamma_temperature_scales_the_matching_latent(monkeypatch): + schedule = tuple(.4 + .1 * index for index in range(len(SP.GAMMAS))) + monkeypatch.setattr(E, "TEMPERATURE_BY_GAMMA", schedule) + latents = np.ones((3, 4), np.float32) + scaled = E._latent_tensor(latents, [0, 2, 6], device="cpu") + assert scaled[:, 0].numpy() == pytest.approx( + [schedule[0], schedule[2], schedule[6]] + ) + + +def test_global_temperature_preserves_original_tensor_arithmetic(monkeypatch): + temperature = .55 + monkeypatch.setattr(E, "TEMPERATURE", temperature) + monkeypatch.setattr( + E, + "TEMPERATURE_BY_GAMMA", + tuple(temperature for _ in SP.GAMMAS), + ) + latents = np.random.default_rng(7).standard_normal( + (3, 4), dtype=np.float32 + ) + expected = temperature * torch.as_tensor(latents) + actual = E._latent_tensor(latents, [0, 2, 6], device="cpu") + assert torch.equal(actual, expected) + + +def test_gamma_temperature_requires_exactly_seven_values(tmp_path): + args = E.build_parser().parse_args([ + "--checkpoints", str(tmp_path / "missing.pt"), + "--labels", "r0", + "--output-dir", str(tmp_path / "out"), + "--temperature-by-gamma", "0.5", "0.6", + ]) + args.m_per_gamma = 1 + with pytest.raises(ValueError, match="seven"): + E.run(args) + + +def test_render_uses_ball_style_validity_name_and_writes_manifest(tmp_path): + records = [] + for round_i in (0, 1, 2): + rows = [ + _compact_row( + episode, + gamma, + validity=min(1.0, .4 + .1 * round_i), + ) + for gamma in SP.GAMMAS + for episode in (1, 2) + ] + summary = E.summarize(rows, seed=round_i + 10) + records.append({ + "label": f"r{round_i}", + "round": round_i, + "cell": {"summary": summary}, + }) + + outputs = E.render(records, str(tmp_path)) + assert {Path(path).suffix for path in outputs} == {".png", ".pdf", ".json"} + assert all(Path(path).stat().st_size > 0 for path in outputs) + manifest_path = next(Path(path) for path in outputs if path.endswith(".json")) + manifest = json.loads(manifest_path.read_text()) + assert "Validity" in manifest["claim"] + assert "V_safe" not in manifest["claim"] + assert [title for _, title, _ in E.PLOT_SPECS] == [ + "Collision rate", + "Validity", + "Min. clearance [m]", + "Time-to-goal [s]", + ] diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_b1_offline_eval_funnel.py b/overnight_run_07_12_sfm/analysis/test_sfm_b1_offline_eval_funnel.py new file mode 100644 index 0000000..11f9069 --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_sfm_b1_offline_eval_funnel.py @@ -0,0 +1,68 @@ +from pathlib import Path +import sys +from types import SimpleNamespace + + +ROOT = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(ROOT)) + +import run_sfm_b1_offline_eval_funnel as FUNNEL # noqa: E402 +import sfm_b1_offline_eval as EVAL # noqa: E402 + + +def _row(name, round_index, *, sr, cr, validity): + return { + "arm": name, + "round": round_index, + "checkpoint": f"/tmp/{name}_{round_index}.pt", + "checkpoint_sha256": f"{name}-{round_index}", + "SR": sr, + "CR": cr, + "timeout": 1.0 - sr - cr, + "Validity": validity, + "clearance": 0.1, + "time_to_goal": 10.0, + } + + +def test_selection_rejects_zero_success_low_collision_collapse(): + r0 = _row("pretrained", 0, sr=0.6, cr=0.4, validity=0.5) + collapsed = _row("a", 2, sr=0.0, cr=0.0, validity=0.9) + viable = _row("b", 1, sr=0.7, cr=0.2, validity=0.7) + selected, contract = FUNNEL.choose_candidates( + [collapsed, viable], r0, top_k=1 + ) + assert selected == [viable] + assert contract["r0_SR_gate"] == 0.6 + assert not contract["fallback_used"] + + +def test_selection_fallback_is_highest_success(): + r0 = _row("pretrained", 0, sr=0.8, cr=0.2, validity=0.5) + first = _row("a", 1, sr=0.4, cr=0.1, validity=0.9) + second = _row("b", 2, sr=0.7, cr=0.3, validity=0.6) + selected, contract = FUNNEL.choose_candidates( + [first, second], r0, top_k=1 + ) + assert selected == [second] + assert contract["fallback_used"] + + +def test_evaluator_artifacts_follow_requested_m(tmp_path): + previous = EVAL.M_PER_GAMMA + try: + EVAL.M_PER_GAMMA = 10 + assert EVAL._artifact_prefix() == "raw_m10_offline" + assert EVAL._status() == "SFM_B1_OFFLINE_RAW_M10_COMPLETE" + finally: + EVAL.M_PER_GAMMA = previous + + +def test_evaluator_rejects_nonpositive_m_before_checkpoint_loading(): + args = SimpleNamespace(m_per_gamma=0) + try: + EVAL.run(args) + except ValueError as error: + assert "--m-per-gamma must be positive" in str(error) + else: + raise AssertionError("nonpositive M must fail") diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_b1_offline_store_replay.py b/overnight_run_07_12_sfm/analysis/test_sfm_b1_offline_store_replay.py new file mode 100644 index 0000000..81b20d5 --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_sfm_b1_offline_store_replay.py @@ -0,0 +1,381 @@ +import json +import math + +import numpy as np +import pytest +import torch + +import sfm_b1_offline_exec as OE +import sfm_b1_offline_replay as OR +import sfm_b1_offline_store as OS + + +def _result(y, *, resolved=True, full_h=True): + if not resolved: + return dict(resolved=False, error="solver") + return dict( + resolved=True, + y=int(y), + taskspace=bool(y), + collision_free=bool(y), + certificate=bool(y), + full_h=bool(full_h), + terminal_step=10 if full_h else 4, + diagnostics={"margin": 0.25}, + ) + + +def _context(shard, *, scenario, gamma, step): + return shard.add_context( + scenario_id=scenario, + gamma=gamma, + step=step, + state=np.zeros(4, np.float32), + hp10=np.zeros((10, 16, 12), np.float32), + low5=np.zeros(5, np.float32), + hist=np.zeros((16, 2), np.float32), + ped_xy=np.zeros((1, 2), np.float32), + ped_vel=np.zeros((1, 2), np.float32), + ) + + +def _add_window(shard, *, scenario, gamma, step, y): + context_id = _context( + shard, scenario=scenario, gamma=gamma, step=step, + ) + controls = np.full((10, 2), scenario + step / 100.0, np.float32) + x0 = np.full(20, scenario - step / 100.0, np.float32) + shard.add_executed_window( + context_id, + controls, + x0, + _result(y), + execution_source="selected_B" if y else "raw_continuation", + nvp_context=not bool(y), + candidate_id=0 if y else None, + acquisition_step=0 if y else None, + sigma=0.4 if y else None, + hp_margin=0.2, + mode="U" if scenario % 2 else "R", + ) + return context_id + + +def _mixed_shard(positive=7, negative=3): + shard = OS.ExecutedRoundShard(1) + gammas = (0.1, 0.2, 0.3, 0.4, 0.5, 0.7, 1.0) + for index in range(positive + negative): + _add_window( + shard, + scenario=100 + index, + gamma=gammas[index % len(gammas)], + step=index, + y=index < positive, + ) + return shard + + +class _TinyPolicy(torch.nn.Module): + """Minimal policy surface needed by the offline replay implementation.""" + + def __init__(self): + super().__init__() + self.enc_grid = torch.nn.Linear(1, 1, bias=False) + self.head = torch.nn.Linear(20, 20, bias=False) + self.d = 20 + self.u_max = 2.0 + + def ctx_from(self, grid, low, hist): + del grid, hist + return low[:, :1] + + def forward(self, value, tau, context): + del tau, context + return self.head(value) + + def cfm_loss(self, controls, context, weights=None): + del context + value = controls.reshape(len(controls), self.d) / self.u_max + per = (self.head(value) - value).square().mean(dim=1) + if weights is None: + return per.mean() + return (per * weights).sum() / weights.sum() + + def module_groups(self): + return {"E_g": self.enc_grid, "head": self.head} + + +def _trainable(policy): + for parameter in policy.parameters(): + parameter.requires_grad_(True) + for parameter in policy.enc_grid.parameters(): + parameter.requires_grad_(False) + return [ + parameter for parameter in policy.parameters() + if parameter.requires_grad + ] + + +def test_executed_store_has_one_window_per_context_and_exact_partition(tmp_path): + shard = OS.ExecutedRoundShard(3) + positive_context = _add_window( + shard, scenario=11, gamma=0.1, step=2, y=1, + ) + _add_window(shard, scenario=12, gamma=1.0, step=3, y=0) + + with pytest.raises(ValueError, match="at most one executed window"): + shard.add_executed_window( + positive_context, + np.zeros((10, 2), np.float32), + np.zeros(20, np.float32), + _result(1), + execution_source="selected_B", + nvp_context=False, + ) + with pytest.raises(ValueError, match="exact full-H=10"): + context_id = _context( + shard, scenario=13, gamma=0.5, step=4, + ) + shard.add_executed_window( + context_id, + np.zeros((10, 2), np.float32), + np.zeros(20, np.float32), + _result(1, full_h=False), + execution_source="selected_B", + nvp_context=False, + ) + with pytest.raises(ValueError, match=r"finite controls \[10,2\]"): + context_id = _context( + shard, scenario=14, gamma=0.5, step=5, + ) + shard.add_executed_window( + context_id, + np.zeros((9, 2), np.float32), + np.zeros(20, np.float32), + _result(1), + execution_source="selected_B", + nvp_context=False, + ) + with pytest.raises(ValueError, match=r"original x0 \[20\]"): + context_id = _context( + shard, scenario=15, gamma=0.5, step=6, + ) + shard.add_executed_window( + context_id, + np.zeros((10, 2), np.float32), + np.zeros(19, np.float32), + _result(1), + execution_source="selected_B", + nvp_context=False, + ) + + assert len(shard.D) == 2 + assert [row["y"] for row in shard.Dplus] == [1] + assert [row["y"] for row in shard.Dminus] == [0] + assert {row["window_id"] for row in shard.D} == {0, 1} + assert {row["context_id"] for row in shard.D} == {0, 1} + assert shard.validate() == { + "round": 3, + "contexts": 5, + "D": 2, + "Dplus": 1, + "Dminus": 1, + "errors": 0, + "unresolved_contexts": 3, + } + + path = tmp_path / "round_003.pt" + manifest = shard.save(path) + assert manifest["D"] == manifest["Dplus"] + manifest["Dminus"] == 2 + with open(str(path) + ".COMPLETE.json") as stream: + marker = json.load(stream) + assert marker["status"] == "OFFLINE_EXECUTED_ROUND_SHARD_COMPLETE" + assert marker["sha256"] == OS.sha256_file(path) + + restored = OS.ExecutedRoundShard.load(path) + assert restored.validate() == shard.validate() + assert [row["execution_source"] for row in restored.D] == [ + "selected_B", "raw_continuation", + ] + np.testing.assert_array_equal( + restored.Dminus[0]["controls"], shard.Dminus[0]["controls"], + ) + np.testing.assert_array_equal( + restored.Dminus[0]["x0"], shard.Dminus[0]["x0"], + ) + + +def test_phi_s_from_x0_is_the_exact_noised_representation(): + from flow_policy import FlowPolicy + + torch.manual_seed(8) + policy = FlowPolicy(T=10, ctx_dim=3, width=12, depth=2, u_max=2.0) + controls = torch.randn(2, 10, 2) + context = torch.randn(2, 3) + x0 = torch.randn(2, 20) + s = 0.9 + + actual = policy.phi_s_from_x0(controls, context, x0, s=s) + x1 = (controls / policy.u_max).reshape(2, 20) + expected = policy.features( + (1 - s) * x0 + s * x1, + torch.full((2,), s), + context, + ) + torch.testing.assert_close(actual, expected) + assert not torch.equal( + actual, + policy.phi_s_from_x0(controls, context, x0.flip(0), s=s), + ) + + +def test_stratified_batches_are_deterministic_and_exact_once(): + shard = _mixed_shard() + left, left_positive, left_negative = OR.stratified_batches( + shard, batch=4, seed=73, + ) + right, _, _ = OR.stratified_batches(shard, batch=4, seed=73) + left_ids = [ + (record[0].round_i, record[1]["window_id"]) + for batch in left for record in batch + ] + right_ids = [ + (record[0].round_i, record[1]["window_id"]) + for batch in right for record in batch + ] + assert left_ids == right_ids + assert len(left_ids) == len(set(left_ids)) == len(shard.D) + assert len(left_positive) == len(shard.Dplus) == 7 + assert len(left_negative) == len(shard.Dminus) == 3 + assert len(left) == math.ceil(len(shard.D) / 4) + assert all(any(record[1]["y"] == 1 for record in batch) for batch in left) + + +def test_stratified_batches_prevent_negative_only_tail(): + shard = _mixed_shard(positive=5, negative=8) + batches, positives, negatives = OR.stratified_batches( + shard, batch=4, seed=74, + ) + assert len(batches) == math.ceil((len(positives) + len(negatives)) / 4) + assert all(len(values) <= 4 for values in batches) + assert all(any(record[1]["y"] == 1 for record in values) for values in batches) + assert sum(len(values) for values in batches) == len(shard.D) + + +def test_stratified_batches_fail_when_sign_safe_partition_is_impossible(): + shard = _mixed_shard(positive=1, negative=8) + with pytest.raises(RuntimeError, match="positive in every"): + OR.stratified_batches(shard, batch=4, seed=75) + + +@pytest.mark.parametrize("exposure_epochs", (1, 10, 100)) +def test_replay_exact_exposure_counts_and_adam_step_formula(exposure_epochs): + torch.manual_seed(14) + shard = _mixed_shard() + policy = _TinyPolicy() + optimizer = torch.optim.Adam(_trainable(policy), lr=1.0e-4) + report = OR.replay( + policy, + optimizer, + shard, + alpha=0.01, + exposure_epochs=exposure_epochs, + batch=4, + device="cpu", + seed=101, + ) + steps_per_epoch = math.ceil(len(shard.D) / 4) + assert report["batches_per_epoch"] == steps_per_epoch + assert report["optimizer_steps"] == steps_per_epoch * exposure_epochs + assert report["positive_total_visits"] == len(shard.Dplus) * exposure_epochs + assert report["negative_total_visits"] == len(shard.Dminus) * exposure_epochs + assert all(row["positive_visits"] == len(shard.Dplus) for row in report["epochs"]) + assert all(row["negative_visits"] == len(shard.Dminus) for row in report["epochs"]) + assert report["exact_once_per_exposure_epoch"] + assert report["negative_used_for_training"] + assert report["visual_encoder_sha_before"] == report["visual_encoder_sha_after"] + + +def test_alpha_zero_retains_and_counts_Dminus_but_never_uses_negative_gradient( + monkeypatch, +): + torch.manual_seed(15) + shard = _mixed_shard() + policy = _TinyPolicy() + optimizer = torch.optim.Adam(_trainable(policy), lr=1.0e-4) + original = OR._weighted_loss + negative_loss_calls = [] + + def audit_weighted_loss(policy, records, mass, population, device): + if records and all(int(record[1]["y"]) == 0 for record in records): + negative_loss_calls.append(len(records)) + return original(policy, records, mass, population, device) + + monkeypatch.setattr(OR, "_weighted_loss", audit_weighted_loss) + report = OR.replay( + policy, + optimizer, + shard, + alpha=0.0, + exposure_epochs=1, + batch=4, + device="cpu", + seed=102, + ) + assert len(shard.Dminus) == report["negative_eligible"] == 3 + assert report["negative_total_visits"] == 3 + assert not report["negative_used_for_training"] + assert negative_loss_calls == [] + assert report["optimizer_steps"] == math.ceil(len(shard.D) / 4) + assert all(row["negative_loss"] is None for row in report["epochs"]) + + +def test_gp_cap_512_has_equal_gamma_quota_and_rotating_extra(): + shard = OS.ExecutedRoundShard(1) + for gamma_index, gamma in enumerate(OE.SP.GAMMAS): + for sample_index in range(74): + _add_window( + shard, + scenario=1_000 + gamma_index, + gamma=gamma, + step=sample_index, + y=1, + ) + + selected_round_2, report_round_2 = OE._gamma_balanced_records( + shard, cap=512, round_i=2, seed=20, + ) + selected_round_3, report_round_3 = OE._gamma_balanced_records( + shard, cap=512, round_i=3, seed=20, + ) + + assert len(selected_round_2) == len(selected_round_3) == 512 + assert report_round_2["quota"] == report_round_3["quota"] == 73 + assert report_round_2["unique"] and report_round_3["unique"] + assert report_round_2["rotating_extra_gamma"] == 0.1 + assert report_round_3["rotating_extra_gamma"] == 0.2 + assert report_round_2["per_gamma"]["0.1"] == 74 + assert report_round_3["per_gamma"]["0.2"] == 74 + for gamma in OE.SP.GAMMAS: + expected_round_2 = 74 if gamma == 0.1 else 73 + expected_round_3 = 74 if gamma == 0.2 else 73 + assert report_round_2["per_gamma"][str(gamma)] == expected_round_2 + assert report_round_3["per_gamma"][str(gamma)] == expected_round_3 + + +def test_gp_quota_fails_closed_instead_of_redistributing_gamma_shortfall(): + shard = OS.ExecutedRoundShard(1) + for gamma_index, gamma in enumerate(OE.SP.GAMMAS): + count = 72 if gamma == 0.1 else 74 + for sample_index in range(count): + _add_window( + shard, + scenario=2_000 + gamma_index, + gamma=gamma, + step=sample_index, + y=1, + ) + with pytest.raises(RuntimeError, match="strict gamma-balanced GP quota"): + OE._gamma_balanced_records( + shard, cap=512, round_i=2, seed=21, + ) diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_b1_protocol.py b/overnight_run_07_12_sfm/analysis/test_sfm_b1_protocol.py index 0501c16..1356aeb 100644 --- a/overnight_run_07_12_sfm/analysis/test_sfm_b1_protocol.py +++ b/overnight_run_07_12_sfm/analysis/test_sfm_b1_protocol.py @@ -17,6 +17,9 @@ def test_frozen_arm_matrix_and_macro_round_ids(): assert len(first) == len(set(first)) == 8 assert set(first).isdisjoint(second) assert 8 * len(P.GAMMAS) == 56 + X.ArmConfig( + name="diagnostic", selector="balanced_rank", alpha=0.0, + ).validate() def test_expansion_has_no_forbidden_legacy_or_expert_path(): diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_b1_r2_aggregate.py b/overnight_run_07_12_sfm/analysis/test_sfm_b1_r2_aggregate.py new file mode 100644 index 0000000..cd0f3f6 --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_sfm_b1_r2_aggregate.py @@ -0,0 +1,43 @@ +import pytest + +import sfm_b1_r2_aggregate as A + + +def test_arm_grid_names_are_unique(): + names = { + A.arm_name(alpha, epochs) + for alpha in A.ALPHAS + for epochs in A.REPLAY_EPOCHS + } + assert len(names) == 9 + assert "margin_alpha0p1_epochs100" in names + + +def test_selection_is_safety_first(): + safe = { + "CR": 0.2, "SR": 0.7, "clearance": 0.1, "time": 10.0, + "round": 2, "alpha": 0.1, "replay_epochs": 100, + } + fast = { + "CR": 0.3, "SR": 0.9, "clearance": 0.2, "time": 5.0, + "round": 1, "alpha": 0.0, "replay_epochs": 1, + } + assert min((safe, fast), key=A._post_expansion_key) is safe + + +def test_paired_cluster_delta_respects_episode_pairing(): + baseline, candidate = [], [] + for episode in (1, 2): + for gamma in (0.1, 1.0): + baseline.append({ + "episode": episode, "gamma": gamma, "collision": episode == 1, + }) + candidate.append({ + "episode": episode, "gamma": gamma, "collision": False, + }) + value = A.paired_cluster_delta( + baseline, candidate, "collision", seed=1, draws=1000, + ) + assert value["estimate"] == pytest.approx(-0.5) + assert value["paired_scenarios"] == 2 + assert value["paired_gamma_cells"] == 4 diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_b1_r2_alpha_replay.py b/overnight_run_07_12_sfm/analysis/test_sfm_b1_r2_alpha_replay.py new file mode 100644 index 0000000..d083b7f --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_sfm_b1_r2_alpha_replay.py @@ -0,0 +1,183 @@ +import copy + +import numpy as np +import pytest +import torch + +import grid_policy_sfm as GPS +import sfm_b1_r2_alpha_replay as R +import sfm_b1_store as BS + + +def _verifier_result(y): + return dict( + resolved=True, y=int(y), taskspace=bool(y), + collision_free=bool(y), certificate=bool(y), full_h=True, + terminal_step=10, train_eligible=bool(y), + segment=np.zeros((11, 2), np.float32), + pedestrian_prediction=np.zeros((11, 1, 2), np.float32), + diagnostics={}, + ) + + +def _recent(tmp_path): + rng = np.random.RandomState(71) + recent = BS.RecentRounds(tmp_path) + for round_i in (1, 2): + shard = BS.RoundShard(round_i) + for gamma in (0.1, 1.0): + context_id = shard.add_context( + scenario_id=100 + round_i, gamma=gamma, step=0, + state=np.zeros(4, np.float32), + hp10=rng.randn(10, 16, 12).astype(np.float32), + low5=rng.randn(5).astype(np.float32), + hist=rng.randn(16, 2).astype(np.float32), + ped_xy=np.zeros((1, 2), np.float32), + ped_vel=np.zeros((1, 2), np.float32), + ) + for candidate_id, y in enumerate((1, 1, 0)): + shard.add_resolved_query( + context_id, candidate_id, + rng.randn(10, 2).astype(np.float32), + sigma=0.4, result=_verifier_result(y), + acquisition_step=candidate_id, + ) + recent.append_and_save(shard) + return recent + + +def _negative_only_recent(tmp_path): + rng = np.random.RandomState(81) + recent = BS.RecentRounds(tmp_path) + shard = BS.RoundShard(1) + context_id = shard.add_context( + scenario_id=300, gamma=0.5, step=0, + state=np.zeros(4, np.float32), + hp10=rng.randn(10, 16, 12).astype(np.float32), + low5=rng.randn(5).astype(np.float32), + hist=rng.randn(16, 2).astype(np.float32), + ped_xy=np.zeros((1, 2), np.float32), + ped_vel=np.zeros((1, 2), np.float32), + ) + for candidate_id in range(2): + shard.add_resolved_query( + context_id, candidate_id, rng.randn(10, 2).astype(np.float32), + sigma=0.5, result=_verifier_result(0), + acquisition_step=candidate_id, + ) + recent.append_and_save(shard) + return recent + + +def test_declared_grid_is_exact_and_other_knobs_fail_closed(): + names = set() + for alpha in R.ALPHAS: + for epochs in R.REPLAY_EPOCHS: + cfg = R.ExperimentConfig(alpha=alpha, replay_epochs=epochs) + assert cfg.validate() is cfg + names.add(cfg.arm_name) + assert len(names) == 9 + with pytest.raises(ValueError): + R.ExperimentConfig(alpha=0.001, replay_epochs=1).validate() + with pytest.raises(ValueError): + R.ExperimentConfig(alpha=0.0, replay_epochs=4).validate() + with pytest.raises(ValueError): + R.ExperimentConfig(alpha=0.0, replay_epochs=1, lr=1e-5).validate() + required_by_gather = ( + "K", "B", "T", "H", "nfe", "temp", "phi_s", "selector", + ) + assert all(hasattr(R.ExperimentConfig(0.0, 1), key) for key in required_by_gather) + + +def test_fixed_probe_is_deterministic(tmp_path): + recent = _recent(tmp_path) + policy = GPS.build_sfm_policy(width=16, res_dropout=0.0) + positives = recent.positive_records() + left = R._fixed_probe_loss( + policy, positives, batch=3, device="cpu", seed=123, + ) + right = R._fixed_probe_loss( + policy, positives, batch=3, device="cpu", seed=123, + ) + assert left == right + + +def test_alpha_zero_never_reads_negative_and_one_epoch_is_complete(tmp_path, monkeypatch): + torch.manual_seed(17) + recent = _recent(tmp_path) + policy = GPS.build_sfm_policy(width=16, res_dropout=0.0) + BS.configure_expansion_trainability(policy) + optimizer = torch.optim.Adam( + [parameter for parameter in policy.parameters() if parameter.requires_grad], + lr=R.LEARNING_RATE, + ) + monkeypatch.setattr( + recent, "negative_records", + lambda: (_ for _ in ()).throw(AssertionError("alpha=0 read D-")), + ) + cfg = R.ExperimentConfig(alpha=0.0, replay_epochs=1) + report = R.repeat_complete_replay( + policy, optimizer, recent, cfg, device="cpu", round_i=1, + ) + assert report["optimizer_steps"] == 1 + assert report["positive_total_visits"] == report["positive_eligible"] + assert report["negative_eligible"] == 0 + assert report["negative_total_visits"] == 0 + assert not report["negative_used_for_training"] + assert report["fixed_probe"]["before"]["negative"] is None + assert report["fixed_probe"]["after"]["negative"] is None + assert report["epochs"][0]["positive_coverage"]["exact_once"] + assert report["visual_encoder_sha_before"] == report["visual_encoder_sha_after"] + assert report["module_relative_parameter_drift"]["E_g"] == 0.0 + + +def test_signed_replay_repeats_complete_support_per_epoch(tmp_path): + torch.manual_seed(19) + recent = _recent(tmp_path) + policy = GPS.build_sfm_policy(width=16, res_dropout=0.0) + initial = copy.deepcopy(policy) + BS.configure_expansion_trainability(policy) + optimizer = torch.optim.Adam( + [parameter for parameter in policy.parameters() if parameter.requires_grad], + lr=R.LEARNING_RATE, + ) + cfg = R.ExperimentConfig(alpha=0.1, replay_epochs=10) + report = R.repeat_complete_replay( + policy, optimizer, recent, cfg, device="cpu", round_i=2, + ) + assert report["optimizer_steps"] == 10 + assert report["positive_total_visits"] == 10 * report["positive_eligible"] + assert report["negative_total_visits"] == 10 * report["negative_eligible"] + assert all(row["positive_coverage"]["exact_once"] for row in report["epochs"]) + assert all(row["negative_coverage"]["exact_once"] for row in report["epochs"]) + assert all(row["gradient_cosine"] is not None for row in report["epochs"]) + assert report["module_relative_parameter_drift"]["E_g"] == 0.0 + assert report["visual_encoder_sha_before"] == report["visual_encoder_sha_after"] + assert any( + not torch.equal(value, initial.state_dict()[name]) + for name, value in policy.state_dict().items() + if not name.startswith("enc_grid.") + ) + + +def test_signed_negative_only_audits_every_record_and_takes_no_step(tmp_path): + torch.manual_seed(29) + recent = _negative_only_recent(tmp_path) + policy = GPS.build_sfm_policy(width=16, res_dropout=0.0) + BS.configure_expansion_trainability(policy) + initial = copy.deepcopy(policy.state_dict()) + optimizer = torch.optim.Adam( + [parameter for parameter in policy.parameters() if parameter.requires_grad], + lr=R.LEARNING_RATE, + ) + cfg = R.ExperimentConfig(alpha=0.1, replay_epochs=10) + report = R.repeat_complete_replay( + policy, optimizer, recent, cfg, device="cpu", round_i=1, + ) + assert report["positive_eligible"] == 0 + assert report["optimizer_steps"] == 0 + assert report["negative_total_visits"] == 10 * report["negative_eligible"] + assert all(row["path"] == "signed_no_positive" for row in report["epochs"]) + assert all(row["negative_coverage"]["exact_once"] for row in report["epochs"]) + for name, value in policy.state_dict().items(): + assert torch.equal(value, initial[name]) diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_b1_r2_eval.py b/overnight_run_07_12_sfm/analysis/test_sfm_b1_r2_eval.py new file mode 100644 index 0000000..c3cf09f --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_sfm_b1_r2_eval.py @@ -0,0 +1,135 @@ +from __future__ import annotations + +from pathlib import Path + +import numpy as np +import pytest + +import sfm_b1_r2_eval as E +import sfm_protocol as SP + + +def _row(episode, gamma, *, status, clearance=None, time=None, v_safe=False): + return { + "episode": int(episode), + "gamma": float(gamma), + "status": status, + "success": status == "success", + "collision": status == "collision", + "timeout": status == "timeout", + "successful_clearance": clearance, + "time_to_goal": time, + "v_safe": bool(v_safe), + "verifier_errors": 0, + "certified_windows": 10, + } + + +def test_archive_reference_is_separate_and_m50_bank_is_disjoint(): + E._assert_disjoint_from_archive(260_000) + with pytest.raises(ValueError, match="disjoint"): + E._assert_disjoint_from_archive(250_050) + assert E.ARCHIVED_M100_REFERENCE["M_per_gamma"] == 100 + assert E.M_PER_GAMMA == 50 + assert "not_a_curve_point" in E.ARCHIVED_M100_REFERENCE["role"] + + +def test_checkpoint_specs_require_unique_increasing_round_labels(tmp_path): + checkpoints = [] + for name in ("a.pt", "b.pt"): + path = tmp_path / name + path.write_bytes(b"x") + checkpoints.append(str(path)) + specs = E._checkpoint_specs(checkpoints, ["r0", "r1"]) + assert [spec["round"] for spec in specs] == [0, 1] + with pytest.raises(ValueError, match="form"): + E._checkpoint_specs(checkpoints, ["pretrained", "r1"]) + with pytest.raises(ValueError, match="increasing"): + E._checkpoint_specs(checkpoints, ["r1", "r0"]) + + +def test_summarize_uses_actual_collision_and_success_only_continuous_metrics(): + rows = [] + for gamma in SP.GAMMAS: + rows.extend([ + _row( + 260_000, gamma, status="success", clearance=0.2, + time=9.0, v_safe=True, + ), + _row( + 260_001, gamma, status="timeout", clearance=None, + time=None, v_safe=True, + ), + _row( + 260_002, gamma, status="collision", clearance=None, + time=None, v_safe=False, + ), + ]) + summary = E.summarize(rows, seed=7) + pooled = summary["pooled"] + assert pooled["SR"] == pytest.approx(1 / 3) + assert pooled["CR"] == pytest.approx(1 / 3) + assert pooled["timeout"] == pytest.approx(1 / 3) + assert pooled["V_safe"] == pytest.approx(2 / 3) + assert pooled["successful_clearance"]["mean"] == pytest.approx(0.2) + assert pooled["successful_clearance"]["n"] == len(SP.GAMMAS) + assert pooled["successful_time_to_goal"]["mean"] == pytest.approx(9.0) + assert pooled["successful_time_to_goal"]["n"] == len(SP.GAMMAS) + + +def test_noise_bank_is_deterministic_and_checkpoint_common(): + first, first_meta = E._noise_bank(ep0=260_000, d=20, seed=123) + second, second_meta = E._noise_bank(ep0=260_000, d=20, seed=123) + assert first.shape == (len(SP.GAMMAS), 50, E.T, 20) + assert first.dtype == np.float32 + assert np.array_equal(first, second) + assert first_meta == second_meta + assert first_meta["temperature"] == 1.0 + assert first_meta["NFE"] == 8 + + +def test_render_writes_paper_style_png_and_pdf(tmp_path): + records = [] + for round_i in (0, 1, 2): + per_gamma = {} + rows = [] + for gamma in SP.GAMMAS: + cell_rows = [ + _row( + 260_000, gamma, status="success", + clearance=0.1 + round_i * 0.01, + time=9.0 + round_i, v_safe=True, + ), + _row( + 260_001, gamma, status="collision", + clearance=None, time=None, v_safe=False, + ), + ] + rows.extend(cell_rows) + per_gamma[str(gamma)] = E._summarize_one(cell_rows, round_i + 1) + pooled = E._summarize_one(rows, round_i + 100) + for metric, key in ( + ("SR", "success"), ("CR", "collision"), + ("timeout", "timeout"), ("V_safe", "v_safe"), + ): + pooled[f"{metric}_cluster_bootstrap95"] = ( + E._cluster_bootstrap_interval(rows, key, seed=round_i + 200) + ) + pooled["successful_clearance"]["cluster_bootstrap95"] = ( + E._cluster_bootstrap_interval( + rows, "successful_clearance", seed=round_i + 300 + ) + ) + pooled["successful_time_to_goal"]["cluster_bootstrap95"] = ( + E._cluster_bootstrap_interval( + rows, "time_to_goal", seed=round_i + 301 + ) + ) + records.append({ + "label": f"r{round_i}", + "round": round_i, + "cell": {"summary": {"pooled": pooled, "per_gamma": per_gamma}}, + }) + outputs = E.render(records, str(tmp_path)) + assert {Path(path).suffix for path in outputs} == {".png", ".pdf"} + assert all(Path(path).stat().st_size > 0 for path in outputs) diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_b1_teacher_branch_viz.py b/overnight_run_07_12_sfm/analysis/test_sfm_b1_teacher_branch_viz.py new file mode 100644 index 0000000..13cdf95 --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_sfm_b1_teacher_branch_viz.py @@ -0,0 +1,134 @@ +import os + +import matplotlib.pyplot as plt +import numpy as np +import pytest +import torch + +import sfm_b1_offline_store as OS +import sfm_b1_teacher_branch_viz as V +import sfm_scene as SS + + +def _result(y): + return dict( + resolved=True, y=int(y), full_h=True, terminal_step=10, + taskspace=bool(y), collision_free=bool(y), certificate=bool(y), + diagnostics=dict(slack=0.1), + ) + + +def _shard(path): + shard = OS.ExecutedRoundShard(1) + for scenario in (10, 11, 12): + for gamma in SS.GAMMAS: + state = np.array([0.2, 0.3, 0.0, 0.0], np.float32) + ped_xy = np.array([[1.0, 1.0]], np.float32) + ped_vel = np.array([[0.1, 0.0]], np.float32) + context_id = shard.add_context( + scenario_id=scenario, gamma=gamma, step=0, state=state, + hp10=np.zeros((10, 16, 12), np.float32), + low5=np.zeros(5, np.float32), + hist=np.zeros((16, 2), np.float32), + ped_xy=ped_xy, ped_vel=ped_vel, + ) + y = int(gamma >= 0.3) + shard.add_executed_window( + context_id, np.zeros((10, 2), np.float32), + np.zeros(20, np.float32), _result(y), + execution_source="unit_test", nvp_context=False, + ) + shard.save(path) + return shard + + +def _teacher(path, shard_path, shard, *, mismatch=False, forbidden=False): + context = shard.contexts[0] + snapshot = { + field: ( + np.asarray(context[field]).copy() + if field in ("state", "ped_xy", "ped_vel") + else context[field] + ) + for field in V.CONTEXT_FIELDS + } + if mismatch: + snapshot["state"][0] += 0.1 + record = dict( + teacher_id=0, context_id=0, + controls=np.full((10, 2), 1.5, np.float32), + source=V.TEACHER_SOURCE, + candidate_family="constant_acceleration_escape", + context_snapshot=snapshot, + ) + if forbidden: + record["y"] = 1 + torch.save(dict( + status=V.TEACHER_STATUS, version=1, round=1, + round_shard_sha256=V._sha256(shard_path), + records=[record], + provenance=dict(unit_test=True), + ), path) + + +def test_load_inputs_enforces_exact_context_and_no_safety_label(tmp_path): + shard_path = os.fspath(tmp_path / "round.pt") + teacher_path = os.fspath(tmp_path / "teacher.pt") + shard = _shard(shard_path) + _teacher(teacher_path, shard_path, shard) + loaded, _, by_context, provenance = V.load_inputs( + shard_path, teacher_path, + ) + assert len(loaded.Dplus) == 15 + assert len(loaded.Dminus) == 6 + assert by_context[0][0]["source"] == V.TEACHER_SOURCE + assert provenance["teacher_records"] == 1 + assert provenance["teacher_families"] == { + "constant_acceleration_escape": 1, + } + + _teacher(teacher_path, shard_path, shard, mismatch=True) + with pytest.raises(ValueError, match="does not match"): + V.load_inputs(shard_path, teacher_path) + _teacher(teacher_path, shard_path, shard, forbidden=True) + with pytest.raises(ValueError, match="safety-label fields"): + V.load_inputs(shard_path, teacher_path) + + +def test_draw_cell_uses_purple_teacher_without_relabeling(tmp_path): + shard_path = os.fspath(tmp_path / "round.pt") + teacher_path = os.fspath(tmp_path / "teacher.pt") + shard = _shard(shard_path) + _teacher(teacher_path, shard_path, shard) + _, _, teachers, _ = V.load_inputs(shard_path, teacher_path) + rows = V._lineages(shard)[(10, 0.1)] + figure, axis = plt.subplots() + V.draw_cell(axis, rows, teachers, through_step=0) + colors = [line.get_color() for line in axis.lines] + teacher_lines = [ + line for line in axis.lines if line.get_color() == V.TEACHER_COLOR + ] + assert V.TEACHER_COLOR in colors + assert all(line.get_linestyle() == "--" for line in teacher_lines) + assert any(line.get_color() == "#111111" for line in axis.lines) + plt.close(figure) + + +def test_render_final_png_and_report(tmp_path): + shard_path = os.fspath(tmp_path / "round.pt") + teacher_path = os.fspath(tmp_path / "teacher.pt") + png = os.fspath(tmp_path / "teacher_branches.png") + report_path = os.fspath(tmp_path / "teacher_branches.json") + shard = _shard(shard_path) + _teacher(teacher_path, shard_path, shard) + report = V.render( + shard_path, teacher_path, png, report_path, + scenarios=(10, 11, 12), + ) + assert report["status"] == "SFM_B1_TEACHER_D_BRANCH_VIZ_COMPLETE" + assert report["teacher_counts"] == dict( + records=1, contexts=1, + families={"constant_acceleration_escape": 1}, + ) + assert os.path.getsize(png) > 0 + assert os.path.getsize(report_path) > 0 diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_b1_verifier.py b/overnight_run_07_12_sfm/analysis/test_sfm_b1_verifier.py index cddd370..6e2edb8 100644 --- a/overnight_run_07_12_sfm/analysis/test_sfm_b1_verifier.py +++ b/overnight_run_07_12_sfm/analysis/test_sfm_b1_verifier.py @@ -13,13 +13,24 @@ def test_label_is_only_task_collision_and_moving_certificate(): assert "progress" not in result and "cost" not in result -def test_early_goal_is_prefix_only_and_not_replay_eligible(): +def test_predicted_goal_crossing_still_certifies_full_h(): + state = np.array([5.4, 6.0, 1.0, 0.0], np.float32) + result = M.verify_query( + state, np.zeros((10, 2)), np.zeros((0, 2)), np.zeros((0, 2)), .5 + ) + assert np.min(np.linalg.norm(result["segment"] - M.SS.GOAL[None], axis=1)) < .5 + assert result["resolved"] and result["terminal_step"] == 10 + assert result["full_h"] and result["y"] == 1 and result["train_eligible"] + + +def test_post_goal_tail_violation_rejects_the_full_window(): state = np.array([5.7, 6.0, 2.0, 0.0], np.float32) result = M.verify_query( state, np.zeros((10, 2)), np.zeros((0, 2)), np.zeros((0, 2)), .5 ) - assert result["resolved"] and result["terminal_step"] < 10 - assert not result["full_h"] and not result["train_eligible"] + assert np.min(np.linalg.norm(result["segment"] - M.SS.GOAL[None], axis=1)) < .5 + assert result["resolved"] and result["terminal_step"] == 10 and result["full_h"] + assert result["y"] == 0 and not result["taskspace"] and not result["train_eligible"] def test_worker_contract_has_no_legacy_theta_grid_argument(): @@ -29,3 +40,38 @@ def test_worker_contract_has_no_legacy_theta_grid_argument(): assert (context, candidate) == (3, 7) assert result["diagnostics"]["solver"] == "exact_2d_angular_interval_socp" assert result["diagnostics"]["K_artificial"] == 16 + + +def test_executed_window_api_accepts_terminal_truncation_only(): + state = np.zeros(4, np.float32) + no_pedestrians = np.zeros((0, 2), np.float32) + for horizon in (1, 10): + result = M.verify_executed_window( + state, np.zeros((horizon, 2), np.float32), + no_pedestrians, no_pedestrians, .5, + ) + assert result["resolved"] and result["y"] == 1 + assert result["window_horizon"] == horizon + assert "full_h" not in result + assert "train_eligible" not in result + for horizon in (0, 11): + result = M.verify_executed_window( + state, np.zeros((horizon, 2), np.float32), + no_pedestrians, no_pedestrians, .5, + ) + assert not result["resolved"] + + +def test_full_h_query_contract_remains_exactly_ten(): + no_pedestrians = np.zeros((0, 2), np.float32) + short = M.verify_query( + np.zeros(4), np.zeros((9, 2)), + no_pedestrians, no_pedestrians, .5, + ) + full = M.verify_query( + np.zeros(4), np.zeros((10, 2)), + no_pedestrians, no_pedestrians, .5, + ) + assert not short["resolved"] + assert full["resolved"] and full["full_h"] + assert full["terminal_step"] == 10 and full["train_eligible"] diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_neutral_autonomous_followup.py b/overnight_run_07_12_sfm/analysis/test_sfm_neutral_autonomous_followup.py new file mode 100644 index 0000000..d9439bb --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_sfm_neutral_autonomous_followup.py @@ -0,0 +1,69 @@ +from types import SimpleNamespace + +import json + +import run_sfm_neutral_autonomous_followup as A + + +def test_followup_stops_without_touching_training_when_r50_goal_is_met( + tmp_path, monkeypatch, +): + source = "a" * 40 + monkeypatch.setattr(A.GLOBAL, "_source_gate", lambda expected=None: source) + monkeypatch.setattr(A.GAMMA, "_gpu_inventory", lambda indices: {"devices": []}) + initial = tmp_path / "initial.json" + initial.write_text("{}") + gamma = tmp_path / "gamma.json" + gamma.write_text(json.dumps({ + "status": A.GAMMA.STATUS, + "ci_clean_four_metric_win": True, + "objective_achieved": True, + "source_commit": source, + "initial_delivery": str(initial), + "initial_delivery_sha256": A.GLOBAL.FUNNEL.sha256_file(initial), + })) + output = tmp_path / "output" + result = A.run(SimpleNamespace( + gamma_delivery=str(gamma), + output_dir=str(output), + poll_seconds=1, + training_gpu=1, + workers=2, + )) + assert result["action"] == "STOP_GOAL_ACHIEVED_AT_R50_OR_EARLIER" + assert (output / "DELIVERY_COMPLETE.json").is_file() + assert not (output / "r100_training").exists() + + +def test_followup_does_not_stop_on_ci_only_without_final_gates(tmp_path, monkeypatch): + source = "a" * 40 + monkeypatch.setattr(A.GLOBAL, "_source_gate", lambda expected=None: source) + monkeypatch.setattr(A.GAMMA, "_gpu_inventory", lambda indices: {"devices": []}) + initial = tmp_path / "missing-initial.json" + initial.write_text("{}") + gamma = tmp_path / "gamma.json" + payload = { + "status": A.GAMMA.STATUS, + "ci_clean_four_metric_win": True, + "objective_achieved": False, + "source_commit": source, + "initial_delivery": str(initial), + "initial_delivery_sha256": A.GLOBAL.FUNNEL.sha256_file(initial), + } + gamma.write_text(json.dumps(payload)) + + def fake_read(path): + if str(path) == str(gamma): + return payload + raise FileNotFoundError(path) + + monkeypatch.setattr(A.GLOBAL, "_read", fake_read) + try: + A.run(SimpleNamespace( + gamma_delivery=str(gamma), output_dir=str(tmp_path / "output"), + poll_seconds=1, training_gpu=1, workers=2, + )) + except FileNotFoundError: + pass + else: + raise AssertionError("CI-only result incorrectly stopped the follow-up") diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_neutral_creative_sanity.py b/overnight_run_07_12_sfm/analysis/test_sfm_neutral_creative_sanity.py new file mode 100644 index 0000000..10f38ce --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_sfm_neutral_creative_sanity.py @@ -0,0 +1,377 @@ +import json +from pathlib import Path + +import numpy as np +import pytest +import torch + +import run_sfm_neutral_autonomous_followup as AUTO +import run_sfm_neutral_creative_sanity as C +import run_sfm_neutral_gamma_temperature as GAMMA +import sfm_b1_cost as COST +import sfm_b1_neutral_multiround as TRAIN +import sfm_protocol as SP + + +def _write(path, value): + path.write_text(json.dumps(value)) + return path + + +def test_trigger_requires_authenticated_failed_autonomous_delivery(tmp_path): + expected_source = "a" * 40 + training_path = _write(tmp_path / "training.json", {"status": "training"}) + global_path = _write(tmp_path / "global.json", {"status": "global"}) + gamma_path = _write(tmp_path / "gamma.json", { + "status": GAMMA.STATUS, + "ci_clean_four_metric_win": True, + "objective_achieved": False, + }) + delivery = { + "status": AUTO.STATUS, + "action": "CREATIVE_SANITY_REQUIRED", + "ci_clean_four_metric_win": True, + "objective_achieved": False, + "source_commit": expected_source, + "selected_arm": "lr1em5_s04", + "r100_training_delivery": str(training_path), + "r100_training_delivery_sha256": C._sha256(training_path), + "r100_global_delivery": str(global_path), + "r100_global_delivery_sha256": C._sha256(global_path), + "r100_gamma_delivery": str(gamma_path), + "r100_gamma_delivery_sha256": C._sha256(gamma_path), + } + delivery_path = _write(tmp_path / "delivery.json", delivery) + trigger = { + "status": AUTO.CREATIVE_TRIGGER_STATUS, + "action": "CREATIVE_SANITY_REQUIRED", + "source_commit": expected_source, + "autonomous_delivery": str(delivery_path), + "autonomous_delivery_sha256": C._sha256(delivery_path), + "selected_arm": delivery["selected_arm"], + "r100_training_delivery": delivery["r100_training_delivery"], + "r100_training_delivery_sha256": delivery[ + "r100_training_delivery_sha256" + ], + "r100_global_delivery": delivery["r100_global_delivery"], + "r100_global_delivery_sha256": delivery[ + "r100_global_delivery_sha256" + ], + "r100_gamma_delivery": delivery["r100_gamma_delivery"], + "r100_gamma_delivery_sha256": delivery[ + "r100_gamma_delivery_sha256" + ], + } + result = C._validate_trigger( + tmp_path / "trigger.json", trigger, + expected_source=expected_source, + ) + assert result["delivery"]["selected_arm"] == "lr1em5_s04" + + trigger["autonomous_delivery_sha256"] = "0" * 64 + with pytest.raises(RuntimeError, match="digest mismatch"): + C._validate_trigger( + tmp_path / "trigger.json", trigger, + expected_source=expected_source, + ) + + +def test_new_evaluation_banks_are_disjoint_from_training_and_prior_banks(): + new = [C._bank_from(530000, 10, "B"), C._bank_from(540000, 10, "A")] + C._validate_new_banks( + new, [C._bank_from(520000, 50, "prior")], {260000, 260001}, + ) + with pytest.raises(RuntimeError, match="overlaps prior"): + C._validate_new_banks( + [C._bank_from(520049, 10, "B")], + [C._bank_from(520000, 50, "prior")], set(), + ) + with pytest.raises(RuntimeError, match="training scenarios"): + C._validate_new_banks( + [C._bank_from(260000, 10, "B")], [], {260003}, + ) + + +def _phase(round_i, phase, *, SR, CR, timeout, validity): + return { + "round": round_i, + "phase": phase, + "pooled": { + "SR": SR, "CR": CR, "timeout": timeout, + "Validity": validity, + }, + } + + +def test_stage_b_retains_d0_unless_disjoint_overwrite_is_clear(): + beneficial = [] + harmful = [] + for round_i in (10, 20, 30): + beneficial.extend([ + _phase(round_i, "post_Dplus", SR=.70, CR=.25, timeout=.05, validity=.70), + _phase(round_i, "post_D0", SR=.71, CR=.23, timeout=.06, validity=.72), + ]) + harmful.extend([ + _phase(round_i, "post_Dplus", SR=.75, CR=.22, timeout=.03, validity=.72), + _phase(round_i, "post_D0", SR=.65, CR=.21, timeout=.13, validity=.73), + ]) + assert not C._d0_gate(beneficial)["clear_D0_overwrite"] + assert C._d0_gate(harmful)["clear_D0_overwrite"] + + +def _metric_row(round_i, *, SR, CR, timeout, validity, clearance=.1, time=9.0): + cell = { + "SR": SR, "CR": CR, "timeout": timeout, "Validity": validity, + "clearance": clearance, "time_to_goal": time, + } + return { + "round": round_i, + "checkpoint": f"/checkpoint/r{round_i}.pt", + "checkpoint_sha256": f"sha{round_i}", + "pooled": dict(cell), + "per_gamma": {str(gamma): dict(cell) for gamma in SP.GAMMAS}, + } + + +def test_stage_a_gate_requires_liveness_gain_safety_and_gamma_trend(): + candidates = [_metric_row(0, SR=.6, CR=.35, timeout=.05, validity=.6)] + controls = [] + for round_i in range(1, 6): + controls.append(_metric_row( + round_i, SR=.60, CR=.35, timeout=.05, validity=.60, + )) + candidates.append(_metric_row( + round_i, SR=.66, CR=.36, timeout=.04, validity=.59, + )) + gate = C._selector_gate(candidates, controls, controls[-1]) + assert gate["passed"] + + candidates[-1]["pooled"]["CR"] = .50 + for gamma in candidates[-1]["per_gamma"]: + candidates[-1]["per_gamma"][gamma]["CR"] = .50 + assert any( + not row["eligible"] + for row in C._selector_gate(candidates, controls, controls[-1])[ + "decisions" + ] + ) + + +def test_early_selector_qualification_does_not_require_r100_equal_dose(): + candidate = [_metric_row(0, SR=.6, CR=.35, timeout=.05, validity=.6)] + control = [_metric_row(1, SR=.60, CR=.35, timeout=.05, validity=.60)] + candidate.append(_metric_row( + 1, SR=.66, CR=.35, timeout=.04, validity=.60, + )) + unreachable_r100 = _metric_row( + 100, SR=.99, CR=.0, timeout=.0, validity=.99, + ) + gate = C._selector_gate(candidate, control, unreachable_r100) + assert gate["passed"] + assert not gate["decisions"][0]["noninferior_to_r100_reference"] + + +def test_encoder_stage_stops_on_drift_regression_or_cr(tmp_path): + marker = { + "round": 1, + "encoder_diagnostics": { + "token_cosine": .97, + "relative_parameter_drift": .001, + "cumulative_from_reference": { + "token_cosine": .97, + "token_rms_change": .01, + "relative_parameter_drift": .001, + }, + }, + "updates": { + "Dplus": {"encoder_gradient_norms": [.1]}, + "D0": {"encoder_gradient_norms": [.2]}, + }, + "paired_trigger_probe": { + "Dplus_increment": {"Dplus_regressed": 0}, + }, + "gather": { + "gp_diagnostics": {"kernel_effective_rank": 12.0}, + "acquisition": {"uplift": .01}, + }, + } + marker_path = _write(tmp_path / "round.json", marker) + delivery = {"round_records": [str(marker_path)]} + candidate = [_metric_row(1, SR=.7, CR=.2, timeout=.1, validity=.7)] + control = [_metric_row(1, SR=.7, CR=.2, timeout=.1, validity=.68)] + control_markers = {1: { + "gather": { + "gp_diagnostics": {"kernel_effective_rank": 11.5}, + "acquisition": {"uplift": .009}, + }, + }} + gate = C._encoder_stage_gate( + delivery, candidate, control, control_markers, + ) + assert gate["unsafe"] + assert not gate["continue_to_round5"] + + +def test_encoder_stage_preserves_earlier_eligible_round(tmp_path): + paths = [] + candidates, controls, control_markers = [], [], {} + for round_i, cosine in ((1, .995), (2, .97)): + marker = { + "round": round_i, + "encoder_diagnostics": { + "token_cosine": cosine, + "relative_parameter_drift": .001, + "cumulative_from_reference": { + "token_cosine": cosine, "token_rms_change": .01, + "relative_parameter_drift": .001, + }, + }, + "updates": { + "Dplus": {"encoder_gradient_norms": [.1]}, + "D0": {"encoder_gradient_norms": [.2]}, + }, + "paired_trigger_probe": { + "Dplus_increment": {"Dplus_regressed": 0}, + }, + "gather": { + "gp_diagnostics": {"kernel_effective_rank": 13.0}, + "acquisition": {"uplift": .012}, + }, + } + path = _write(tmp_path / f"round{round_i}.json", marker) + paths.append(str(path)) + candidates.append(_metric_row( + round_i, SR=.7, CR=.2, timeout=.1, validity=.70, + )) + controls.append(_metric_row( + round_i, SR=.7, CR=.2, timeout=.1, validity=.68, + )) + control_markers[round_i] = { + "gather": { + "gp_diagnostics": {"kernel_effective_rank": 11.0}, + "acquisition": {"uplift": .009}, + }, + } + gate = C._encoder_stage_gate( + {"round_records": paths}, candidates, controls, control_markers, + ) + assert gate["eligible_rounds"] == [1] + assert not gate["continue_to_round5"] + + +def test_full_creative_command_keeps_locked_winner_and_m20(): + command = C._training_command( + checkpoint="pre.pt", output="out", name="arm", rounds=100, + scenario_ep0=1, eval_ep0=2, eval_rounds=(0, 2, 100), + lr=1e-5, inner_steps=1, ell=.2, gp_cap=512, + selector="margin", encoder_lr_ratio=0.0, workers=2, + noise_seed=3, sample_seed=4, audit_seed=5, train_seed=6, + probe_seed=7, eval_M=20, + locked_eval={ + "checkpoint": "best.pt", "round": 2, + "checkpoint_sha256": "abc", + }, + ) + assert command[command.index("--eval-M") + 1] == "20" + assert command[command.index("--locked-eval-round") + 1] == "2" + + +def test_full_creative_extension_resumes_selected_run_not_pretrained( + tmp_path, monkeypatch, +): + source_run = tmp_path / "stage_c5" + source_run.mkdir() + marker = _write(source_run / "round5.json", { + "round": 5, "scenarios": [260008, 260009], + }) + source_delivery = { + "rounds": 5, + "checkpoint": "/pretrained.pt", + "round_records": [str(marker)], + "config": { + "name": "encoder_c5", "lr": 1e-5, "inner_steps": 1, + "ell": .2, "gp_cap": 512, "selector": "margin", + "encoder_lr_ratio": .1, "sample_seed": 1, "audit_seed": 2, + "train_seed": 3, "probe_seed": 4, "neutral_replay": True, + }, + } + _write(source_run / "DELIVERY_COMPLETE.json", source_delivery) + common = {"selected_lock": { + "name": "encoder_c5_r3", "training_run": str(source_run), + "round": 3, "checkpoint": "/locked-r3.pt", + "checkpoint_sha256": "locked-sha", + }} + commands = [] + + def fake_run(command, **_kwargs): + commands.append(command) + output = Path(command[command.index("--output-root") + 1]) + output.mkdir(parents=True) + _write(output / "DELIVERY_COMPLETE.json", {"status": "complete"}) + + monkeypatch.setattr(C, "_run", fake_run) + monkeypatch.setattr(C, "_authenticate_lock_membership", lambda lock: {}) + monkeypatch.setattr(C, "_confirm_full_run", lambda **_kwargs: { + "objective_achieved": False, + }) + result = C._extend_common_winner_to_r100( + common, output=tmp_path / "out", eval_ep0=1, confirm_ep0=2, + noise_seed=3, gpus=[1, 3], workers=4, + ) + command = commands[0] + assert command[command.index("--resume-run-root") + 1] == str(source_run) + assert command[command.index("--rounds") + 1] == "100" + eval_rounds = command[command.index("--eval-rounds") + 1] + assert {3, 5, 10, 100}.issubset(set(map(int, eval_rounds.split(",")))) + assert result["source_lock"]["round"] == 3 + + +def test_legacy_control_gp_uplift_is_read_from_authenticated_trace(tmp_path): + trace_path = tmp_path / "trace.pt" + torch.save({ + "status": TRAIN.RA.STATUS, + "protocol": {"acquisition": {"uplift": 0.0125}}, + }, trace_path) + marker = { + "gather": { + "trace_path": str(trace_path), + "trace_sha256": C._sha256(trace_path), + }, + } + fields = C._control_gp_fields(marker) + assert fields == { + "uplift": 0.0125, + "effective_rank": None, + "source": "authenticated_legacy_trace", + } + + +def test_progress_gated_margin_prefers_moving_goal_progress_then_falls_back(): + state = np.zeros(4, np.float32) + stalled = { + "candidate_id": 0, "hp_margin": 10.0, + "controls": np.zeros((10, 2), np.float32), + } + moving = { + "candidate_id": 1, "hp_margin": 1.0, + "controls": np.full((10, 2), 2.0, np.float32), + } + assert COST.select_progress_gated_margin( + [stalled, moving], state=state, + )["candidate_id"] == 1 + assert COST.select_progress_gated_margin( + [stalled], state=state, + )["candidate_id"] == 0 + + +def test_creative_trainability_modes_are_individual_not_combined(): + assert TRAIN.StudyConfig( + name="A", selector="progress_gated_margin", + encoder_lr_ratio=0.0, + ).validate() + assert TRAIN.StudyConfig( + name="C", selector="margin", encoder_lr_ratio=0.1, + ).validate() + with pytest.raises(ValueError, match="encoder_lr_ratio"): + TRAIN.StudyConfig(name="bad", encoder_lr_ratio=0.2).validate() + assert TRAIN.StudyConfig(name="B", neutral_replay=False).validate() diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_neutral_gamma_temperature.py b/overnight_run_07_12_sfm/analysis/test_sfm_neutral_gamma_temperature.py new file mode 100644 index 0000000..dc0d48d --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_sfm_neutral_gamma_temperature.py @@ -0,0 +1,348 @@ +from __future__ import annotations + +import json + +import pytest + +import run_sfm_neutral_gamma_temperature as G + + +def _row(episode, gamma, *, success=True, collision=False, validity=.5, + clearance=.1, time=8.0): + return { + "episode": int(episode), + "gamma": float(gamma), + "success": bool(success), + "collision": bool(collision), + "timeout": bool(not success and not collision), + "validity": float(validity), + "successful_clearance": float(clearance) if success else None, + "time_to_goal": float(time) if success else None, + } + + +def test_point_metrics_exclude_failures_only_from_conditional_metrics(): + rows = [ + _row(1, .1, clearance=.2, time=9), + _row(2, .1, success=False, collision=True), + ] + point = G._point(rows) + assert point["SR"] == .5 + assert point["CR"] == .5 + assert point["clearance"] == .2 + assert point["time_to_goal"] == 9 + + +def test_paired_ci_requires_complete_scenario_clusters(): + rows = [_row(episode, gamma) for episode in range(50) for gamma in G.SP.GAMMAS] + result = G._paired_cluster_ci(rows, rows, seed=1, draws=20) + assert all(value["paired_cluster_95"] == [0.0, 0.0] for value in result.values()) + assert not G._ci_win(result) + + +def test_global_temperature_reference_reuse_is_fail_closed(tmp_path): + initial = tmp_path / "DELIVERY_COMPLETE.json" + initial.write_text("{}") + reference_root = tmp_path / "disjoint_m50" + selection = {} + frozen_records = [] + for method in ("pretrained", "expanded"): + scenario_ids = list(range(470000, 470050)) + rows = [ + _row(episode, gamma) + for gamma in G.SP.GAMMAS for episode in scenario_ids + ] + payload = { + "scene_profile": "double_density_velocity_ood", + "bank": { + "ep0": 470000, "M_per_gamma": 50, + "scenario_ids": scenario_ids, + }, + "noise_bank": { + "seed": 7, "gammas": list(map(float, G.SP.GAMMAS)), + "NFE": 8, "dtype": "float32", + "shape": [7, 50, 180, 20], "sha256": "a" * 64, + "temperature_by_gamma": [.55] * 7, + }, + "temperature": .55, + "records": [{ + "round": 0, + "cell": { + "rows": rows, "checkpoint_sha256": method, + "checkpoint": f"/{method}.pt", + "scene_profile": "double_density_velocity_ood", + "cell_key": f"cell-{method}", + "summary": { + "pooled": { + "SR": 1.0, "CR": 0.0, "timeout": 0.0, + "Validity": {"mean": .5}, + "successful_clearance": {"mean": .1}, + "successful_time_to_goal": {"mean": 8.0}, + }, + "per_gamma": { + str(gamma): { + "SR": 1.0, "CR": 0.0, "timeout": 0.0, + "Validity": {"mean": .5}, + "successful_clearance": {"mean": .1}, + "successful_time_to_goal": {"mean": 8.0}, + } + for gamma in G.SP.GAMMAS + }, + }, + }, + }], + } + destination = reference_root / method / "raw_m50_offline_metrics.json" + destination.parent.mkdir(parents=True) + destination.write_text(json.dumps(payload)) + selection[method] = { + "temperature": .55, + "round": 0, + "checkpoint": f"/{method}.pt", + "checkpoint_sha256": method, + } + payload["records"][0]["cell"]["checkpoint"] = f"/{method}.pt" + destination.write_text(json.dumps(payload)) + frozen_records.append(G.BASE._record_from_cell( + payload, method=method, temperature=.55, + )) + state = { + "banks": {"disjoint_confirmation": { + "ep0": 470000, "M_per_gamma": 50, "noise_seed": 7, + }}, + "final_records": frozen_records, + } + cells, reuse = G._reuse_global_temperature_cells( + initial, state, selection + ) + assert reuse["status"] == ( + "GLOBAL_TEMPERATURE_CALIBRATION_CELLS_REUSED" + ) + assert cells["pretrained"][.55]["temperature"] == .55 + + expanded_path = ( + reference_root / "expanded" / "raw_m50_offline_metrics.json" + ) + expanded = json.loads(expanded_path.read_text()) + expanded["noise_bank"]["sha256"] = "b" * 64 + expanded_path.write_text(json.dumps(expanded)) + with pytest.raises(RuntimeError, match="do not share CRN bytes"): + G._reuse_global_temperature_cells(initial, state, selection) + expanded["noise_bank"]["sha256"] = "a" * 64 + expanded_path.write_text(json.dumps(expanded)) + + selection["expanded"]["checkpoint_sha256"] = "wrong" + with pytest.raises(RuntimeError, match="reference contract failed for expanded"): + G._reuse_global_temperature_cells(initial, state, selection) + + +def test_prior_calibration_cells_are_content_authenticated(tmp_path): + bank = {"ep0": 470000, "M_per_gamma": 50, "noise_seed": 17} + methods = { + "pretrained": {"round": 0, "checkpoint_sha256": "pre"}, + "expanded": {"round": 2, "checkpoint_sha256": "exp"}, + } + root = tmp_path / "prior" + for method, selected in methods.items(): + scenario_ids = list(range(470000, 470050)) + rows = [ + _row(episode, gamma) + for gamma in G.SP.GAMMAS for episode in scenario_ids + ] + payload = { + "scene_profile": "double_density_velocity_ood", + "bank": { + "ep0": 470000, "M_per_gamma": 50, + "scenario_ids": scenario_ids, + }, + "noise_bank": { + "seed": 17, "gammas": list(map(float, G.SP.GAMMAS)), + "NFE": 8, "dtype": "float32", + "shape": [7, 50, 180, 20], "sha256": "c" * 64, + }, + "temperature": .7, + "temperature_by_gamma": [.7] * 7, + "records": [{ + "round": selected["round"], + "cell": { + "rows": rows, + "checkpoint_sha256": selected["checkpoint_sha256"], + "scene_profile": "double_density_velocity_ood", + "cell_key": f"{method}-cell", + }, + }], + } + path = ( + root / "calibration_m50" / f"{method}_temp0p7" + / "raw_m50_offline_metrics.json" + ) + path.parent.mkdir(parents=True) + path.write_text(json.dumps(payload)) + cells = {"pretrained": {}, "expanded": {}} + reuse = G._reuse_prior_calibration_cells( + root, methods, bank, [.7], cells, + ) + assert len(reuse["records"]) == 2 + assert not reuse["missing_cells_will_be_run"] + assert set(cells) == {"pretrained", "expanded"} + assert cells["expanded"][.7]["records"][0]["round"] == 2 + + path = ( + root / "calibration_m50" / "expanded_temp0p7" + / "raw_m50_offline_metrics.json" + ) + broken = json.loads(path.read_text()) + broken["records"][0]["round"] = 3 + path.write_text(json.dumps(broken)) + with pytest.raises(RuntimeError, match="prior calibration contract"): + G._reuse_prior_calibration_cells( + root, methods, bank, [.7], + {"pretrained": {}, "expanded": {}}, + ) + + +def test_locked_checkpoint_bytes_and_fresh_cell_are_fail_closed(tmp_path): + checkpoint = tmp_path / "model.pt" + checkpoint.write_bytes(b"frozen") + locked = { + "checkpoint": str(checkpoint), + "checkpoint_sha256": G.BASE.FUNNEL.sha256_file(checkpoint), + "round": 2, + "temperature_by_gamma": [.7] * 7, + } + G._authenticate_method_checkpoints({"expanded": locked}) + rows = [ + _row(episode, gamma) + for gamma in G.SP.GAMMAS for episode in range(480000, 480050) + ] + payload = { + "scene_profile": "double_density_velocity_ood", + "bank": { + "ep0": 480000, "M_per_gamma": 50, + "scenario_ids": list(range(480000, 480050)), + }, + "noise_bank": { + "seed": 9, "gammas": list(map(float, G.SP.GAMMAS)), + "NFE": 8, "dtype": "float32", + "shape": [7, 50, 180, 20], "sha256": "d" * 64, + }, + "temperature_by_gamma": [.7] * 7, + "records": [{ + "round": 2, + "cell": { + "rows": rows, + "checkpoint_sha256": locked["checkpoint_sha256"], + }, + }], + } + assert G._validate_fresh_raw_cell( + payload, locked=locked, ep0=480000, noise_seed=9, + ) == "d" * 64 + + checkpoint.write_bytes(b"mutated") + with pytest.raises(RuntimeError, match="locked expanded checkpoint"): + G._authenticate_method_checkpoints({"expanded": locked}) + payload["records"][0]["cell"]["checkpoint_sha256"] = "0" * 64 + with pytest.raises(RuntimeError, match="fresh raw confirmation"): + G._validate_fresh_raw_cell( + payload, locked=locked, ep0=480000, noise_seed=9, + ) + + +def test_ci_win_requires_all_four_intervals_strictly_favorable(): + result = { + "CR": {"paired_cluster_95": [-.2, -.01]}, + "Validity": {"paired_cluster_95": [.01, .2]}, + "clearance": {"paired_cluster_95": [.001, .02]}, + "time_to_goal": {"paired_cluster_95": [-2.0, -.1]}, + } + assert G._ci_win(result) + result["CR"] = {"paired_cluster_95": [-.2, .01]} + assert not G._ci_win(result) + + +def test_schedule_selection_requires_gamma_trend_before_shortfall(): + def record(name, trend_ok, cr): + pooled = { + "SR": .8, "CR": cr, "timeout": 0.0, + "Validity": .7, "clearance": .2, "time_to_goal": 8.0, + } + rows = {} + for index, gamma in enumerate(G.SP.GAMMAS): + cell = dict(pooled) + cell["clearance"] = (.3 - .01 * index) if trend_ok else (.1 + .05 * index) + cell["time_to_goal"] = (11 - .2 * index) if trend_ok else (7 + 2.0 * index) + cell["Validity"] = .5 + .02 * index + cell["CR"] = .1 + .01 * index + rows[str(gamma)] = cell + return { + "method": name, "round": 1, "pooled": pooled, + "per_gamma": rows, "temperature_by_gamma": [1.0] * 7, + } + + target = {"CR": .1, "Validity": .7, "clearance": .2, "time_to_goal": 8.0} + liveness = {"minimum_SR": .5, "maximum_timeout": .1, "every_gamma_has_success": True} + selected = G._pick( + [record("bad_trend", False, .05), record("good_trend", True, .15)], + target=target, + liveness=liveness, + ) + assert selected["method"] == "good_trend" + + +def test_final_objective_requires_ci_liveness_and_gamma_trend(): + def record(method, *, sr=.8, timeout=0.0, trend_ok=True): + pooled = { + "SR": sr, "CR": .1, "timeout": timeout, + "Validity": .7, "clearance": .2, "time_to_goal": 8.0, + } + per_gamma = {} + for index, gamma in enumerate(G.SP.GAMMAS): + cell = dict(pooled) + cell["clearance"] = (.3 - .01 * index) if trend_ok else (.1 + .05 * index) + cell["time_to_goal"] = (11 - .2 * index) if trend_ok else (7 + 2 * index) + cell["Validity"] = .5 + .02 * index + cell["CR"] = .1 + .01 * index + per_gamma[str(gamma)] = cell + return {"method": method, "pooled": pooled, "per_gamma": per_gamma} + + comparisons = { + name: { + "CR": {"paired_cluster_95": [-.2, -.01]}, + "Validity": {"paired_cluster_95": [.01, .2]}, + "clearance": {"paired_cluster_95": [.001, .02]}, + "time_to_goal": {"paired_cluster_95": [-2.0, -.1]}, + } + for name in ("expanded_minus_pretrained", "expanded_minus_kazuki") + } + good = [record("pretrained"), record("expanded"), record("kazuki_locked")] + assert G._final_objective_gates(good, comparisons)["objective_achieved"] + + bad_trend = [record("pretrained"), record("expanded", trend_ok=False), record("kazuki_locked")] + gates = G._final_objective_gates(bad_trend, comparisons) + assert gates["paired_ci_clean_four_metric_win"] + assert not gates["final_gamma_trend_eligible"] + assert not gates["objective_achieved"] + + bad_liveness = [record("pretrained"), record("expanded", sr=.1, timeout=.8), record("kazuki_locked")] + gates = G._final_objective_gates(bad_liveness, comparisons) + assert not gates["final_liveness_eligible"] + assert not gates["objective_achieved"] + + +def test_one_fully_reversed_trend_family_cannot_hide_in_mean(): + pooled = { + "SR": .8, "CR": .1, "timeout": 0.0, + "Validity": .7, "clearance": .2, "time_to_goal": 8.0, + } + per_gamma = {} + for index, gamma in enumerate(G.SP.GAMMAS): + cell = dict(pooled) + cell["CR"] = .1 + .01 * index + cell["Validity"] = .5 + .02 * index + cell["clearance"] = .3 - .01 * index + cell["time_to_goal"] = 7.0 + 2.0 * index + per_gamma[str(gamma)] = cell + record = {"pooled": pooled, "per_gamma": per_gamma} + assert G.BASE._trend(record)["mean_fraction"] == .75 + assert not G.BASE._trend_eligible(record) diff --git a/overnight_run_07_12_sfm/analysis/test_sfm_neutral_temperature_m50.py b/overnight_run_07_12_sfm/analysis/test_sfm_neutral_temperature_m50.py new file mode 100644 index 0000000..3cb123f --- /dev/null +++ b/overnight_run_07_12_sfm/analysis/test_sfm_neutral_temperature_m50.py @@ -0,0 +1,215 @@ +from __future__ import annotations + +import json +import math +import pytest + +import run_sfm_neutral_temperature_m50 as S +import sfm_protocol as SP + + +def _record(method, *, cr, validity, clearance, time, temperature=1.0): + pooled = { + "SR": 1.0 - cr, + "CR": cr, + "timeout": 0.0, + "Validity": validity, + "clearance": clearance, + "time_to_goal": time, + } + return { + "method": method, + "round": 1, + "temperature": temperature, + "pooled": pooled, + "per_gamma": {str(gamma): dict(pooled) for gamma in SP.GAMMAS}, + } + + +def test_reference_envelope_uses_hardest_metricwise_baseline(): + pretrained = [ + _record("pretrained", cr=.3, validity=.6, clearance=.1, time=9), + _record("pretrained", cr=.2, validity=.5, clearance=.12, time=10), + ] + kazuki = _record( + "kazuki_locked", cr=.1, validity=.4, clearance=.2, time=4, + temperature=None, + ) + assert S._envelope(pretrained, kazuki) == { + "CR": .1, + "Validity": .6, + "clearance": .2, + "time_to_goal": 4, + } + + +def test_four_metric_gate_has_zero_shortfall_only_for_strict_envelope_win(): + target = { + "CR": .2, + "Validity": .6, + "clearance": .1, + "time_to_goal": 9, + } + winner = _record( + "expanded", cr=.1, validity=.7, clearance=.2, time=8 + ) + loser = _record( + "expanded", cr=.1, validity=.7, clearance=.08, time=8 + ) + assert S._shortfalls(winner, target) == { + metric: 0.0 for metric in S.METRICS + } + assert S._shortfalls(loser, target)["clearance"] == pytest.approx(.2) + + +def test_no_success_metrics_can_never_win_selection(): + target = { + "CR": .2, + "Validity": .6, + "clearance": .1, + "time_to_goal": 9, + } + collapsed = _record( + "expanded", cr=0.0, validity=.9, + clearance=float("nan"), time=float("nan"), + ) + shortfall = S._shortfalls(collapsed, target) + assert math.isinf(shortfall["clearance"]) + assert math.isinf(shortfall["time_to_goal"]) + liveness = { + "minimum_SR": .5, + "maximum_timeout": .1, + "every_gamma_has_success": True, + } + assert not S._liveness_eligible(collapsed, liveness) + + +def test_gamma_trend_is_diagnostic_not_per_gamma_temperature_tuning(): + record = _record( + "expanded", cr=.2, validity=.5, clearance=.1, time=9, + temperature=.7, + ) + for index, gamma in enumerate(SP.GAMMAS): + cell = record["per_gamma"][str(gamma)] + cell["CR"] = .1 + .02 * index + cell["Validity"] = .3 + .05 * index + cell["clearance"] = .2 - .01 * index + cell["time_to_goal"] = 12 - .5 * index + trend = S._trend(record) + assert trend["mean_fraction"] == 1.0 + assert record["temperature"] == .7 + + +def test_authenticated_round_records_follow_nested_resume_chain(tmp_path): + lineage = [] + prior_path = None + prior_delivery = None + for round_i in (1, 2, 3): + checkpoint = tmp_path / f"round_{round_i:02d}.pt" + post_positive = tmp_path / f"round_{round_i:02d}_post_positive.pt" + checkpoint.write_bytes(f"checkpoint-{round_i}".encode()) + post_positive.write_bytes(f"positive-{round_i}".encode()) + marker = tmp_path / f"round_{round_i}.json" + marker.write_text(json.dumps({ + "status": "SFM_B1_NEUTRAL_MULTIROUND_ROUND_COMPLETE", + "round": round_i, + "checkpoint": str(checkpoint), + "checkpoint_sha256": S.FUNNEL.sha256_file(checkpoint), + })) + ref = S._round_ref_from_marker(marker) + delivery = { + "status": "SFM_B1_NEUTRAL_MULTIROUND_COMPLETE", + "round_records": [str(marker)], + "round_record_refs": [ref], + } + if prior_delivery is not None: + delivery["resume"] = { + "delivery": str(prior_path), + "delivery_sha256": S.FUNNEL.sha256_file(prior_path), + "round_record_refs": list(lineage), + } + path = tmp_path / f"delivery_{round_i}.json" + path.write_text(json.dumps(delivery)) + lineage.append(ref) + prior_path, prior_delivery = path, delivery + records = S._authenticated_round_records(prior_delivery) + assert [row["round"] for row in records] == [1, 2, 3] + + prior_delivery["round_record_refs"][0]["path"] = str(tmp_path / "wrong") + with pytest.raises(RuntimeError, match="current round-record path"): + S._authenticated_round_records(prior_delivery) + + +def test_legacy_delivery_refs_gain_round_and_checkpoint_identity(tmp_path): + checkpoint = tmp_path / "round_01.pt" + post_positive = tmp_path / "round_01_post_positive.pt" + checkpoint.write_bytes(b"checkpoint") + post_positive.write_bytes(b"positive") + marker = tmp_path / "round.json" + marker.write_text(json.dumps({ + "status": "SFM_B1_NEUTRAL_MULTIROUND_ROUND_COMPLETE", + "round": 1, "checkpoint": str(checkpoint), + "checkpoint_sha256": S.FUNNEL.sha256_file(checkpoint), + })) + delivery = {"round_records": [str(marker)]} + refs = S._delivery_round_refs(delivery) + assert refs[0]["round"] == 1 + assert refs[0]["post_D0_sha256"] == S.FUNNEL.sha256_file(checkpoint) + assert refs[0]["post_Dplus_sha256"] == S.FUNNEL.sha256_file(post_positive) + + checkpoint.write_bytes(b"mutated") + with pytest.raises(RuntimeError, match="checkpoint digest"): + S._authenticated_round_records(delivery) + + +def test_screen_checkpoint_must_belong_to_authenticated_round(tmp_path): + pretrained = tmp_path / "pretrained.pt" + checkpoint = tmp_path / "round_01.pt" + post_positive = tmp_path / "round_01_post_positive.pt" + pretrained.write_bytes(b"pretrained") + checkpoint.write_bytes(b"round-one") + post_positive.write_bytes(b"positive") + marker = tmp_path / "round.json" + marker.write_text(json.dumps({ + "status": "SFM_B1_NEUTRAL_MULTIROUND_ROUND_COMPLETE", + "round": 1, "checkpoint": str(checkpoint), + "checkpoint_sha256": S.FUNNEL.sha256_file(checkpoint), + })) + delivery = { + "checkpoint": str(pretrained), + "checkpoint_sha256": S.FUNNEL.sha256_file(pretrained), + "round_records": [str(marker)], + "round_record_refs": [S._round_ref_from_marker(marker)], + "disjoint_raw_evaluation": {"records": [ + { + "round": 0, "checkpoint": str(pretrained), + "checkpoint_sha256": S.FUNNEL.sha256_file(pretrained), + "pooled": { + "SR": .5, "CR": .5, "timeout": 0, + "Validity": {"mean": .5}, + "successful_clearance": {"mean": .1}, + "successful_time_to_goal": {"mean": 8.0}, + }, + }, { + "round": 1, "checkpoint": str(checkpoint), + "checkpoint_sha256": S.FUNNEL.sha256_file(checkpoint), + "pooled": { + "SR": .5, "CR": .5, "timeout": 0, + "Validity": {"mean": .5}, + "successful_clearance": {"mean": .1}, + "successful_time_to_goal": {"mean": 8.0}, + }, + }]}, + } + rows, _ = S._screen_rows(tmp_path, ["arm"], [delivery]) + assert rows[0]["checkpoint_sha256"] == S.FUNNEL.sha256_file(checkpoint) + legacy = delivery["disjoint_raw_evaluation"]["records"][1] + legacy.pop("checkpoint") + legacy.pop("checkpoint_sha256") + rows, _ = S._screen_rows(tmp_path, ["arm"], [delivery]) + assert rows[0]["checkpoint"] == str(checkpoint.resolve()) + delivery["disjoint_raw_evaluation"]["records"][1][ + "checkpoint_sha256" + ] = "0" * 64 + with pytest.raises(RuntimeError, match="screening checkpoint"): + S._screen_rows(tmp_path, ["arm"], [delivery]) diff --git a/overnight_run_07_12_sfm/claude_afe_driver.py b/overnight_run_07_12_sfm/claude_afe_driver.py new file mode 100644 index 0000000..0d6afeb --- /dev/null +++ b/overnight_run_07_12_sfm/claude_afe_driver.py @@ -0,0 +1,253 @@ +"""Staged AFE-fidelity funnel: phase0 -> screen -> expand -> M50 -> M100.""" +from __future__ import annotations + +import argparse +from concurrent.futures import ProcessPoolExecutor +import copy +import json +import os + +import numpy as np +import torch + +import _paths # noqa: F401 +import claude_afe_fidelity as AF +import claude_continuation as CC +import claude_corrected_distill as CD +import claude_mpc_pool as MP +import grid_policy_sfm as GPS +import sfm_b1_expand as BX +import sfm_b1_offline_eval as OEV +import sfm_b1_offline_store as OS +import sfm_b1_rbf as BR +import sfm_b1_store as BS +import sfm_kazuki as KZ +import sfm_protocol as SP +import sfm_scene as SS + +R1 = ("/data3/research1/claude_sfm_best_recipe_f06e8dd/stageB/" + "B9_margin_origrec_e100/round_01.pt") +B9_SHARD = ("/data3/research1/claude_sfm_best_recipe_f06e8dd/stageB/" + "B9_margin_origrec_e100/round_shards/round_01.pt") +PH0_EP0, M10_EP0, M10_SEED = 460_000, 390_000, 20_260_740 +M50_EP0, M50_SEED, M100_EP0, M100_SEED = 470_000, 20_260_750, 480_000, 20_260_751 +DOSES = (dict(name="d1", lr=1e-5, epochs=4), dict(name="d2", lr=3e-5, epochs=4)) + + +def _phi(policy): + phi = copy.deepcopy(policy).eval() + for p in phi.parameters(): + p.requires_grad_(False) + return phi + + +def _beta(policy, gp, scenarios, gammas, device, ess, executor): + vectors = [] + env = SS.scene_profile("double_density_velocity_ood") + for s in scenarios[:4]: + for g in gammas: + humans = SS.make_humans(int(s), 0, env["n_ped"], + tuple(env["ped_speed_range"])) + ctx = _context0(s, g) + pool = AF.build_pool(policy, ctx, humans, device) + feats = AF.pool_features(policy, ctx, pool["plans"], + pool["x0"], device) + order = torch.as_tensor(AF.keyed_rng( + AF.SEED, "beta", s, g, + ).permutation(len(feats))) + vectors.extend(gp.sequential_score_vectors( + feats, order, AF.B_BUDGET, + )) + beta, achieved = BR.solve_beta(vectors, target=float(ess)) + return float(beta), float(achieved) + + +def _context0(scenario, gamma): + import grid_feats as GF + import sfm_hp_history as HH + env = SS.scene_profile("double_density_velocity_ood") + humans = SS.make_humans(int(scenario), 0, env["n_ped"], + tuple(env["ped_speed_range"])) + ped_xy, ped_vel = SS.collect_humans(humans) + state = np.zeros(4, np.float32) + obstacles = np.concatenate([ + ped_xy, np.full((len(ped_xy), 1), SS.R_PED, np.float32)], axis=1) + hp10 = HH.HpHistory().append(torch.as_tensor(GF.axis_grid( + state[:2], obstacles, 0.0, R=SS.R_SENSE, sensing=SS.R_SENSE))) + return dict(scenario_id=int(scenario), gamma=float(gamma), step=0, + state=state, hp10=hp10.numpy().astype(np.float32), + low5=np.asarray(GF.low5(state, SS.GOAL, gamma), np.float32), + hist=np.zeros((16, 2), np.float32), + ped_xy=ped_xy, ped_vel=ped_vel) + + +def _summ(rows): + n = len(rows) + succ = [r for r in rows if r["status"] == "success"] + cl = [r["min_clearance"] for r in succ] + tt = [r["steps"] * SS.DT for r in succ] + return dict( + n=n, SR=len(succ) / n, + CR=sum(r["status"] == "collision" for r in rows) / n, + NVP=sum(r["status"] == "nvp" for r in rows) / n, + timeout=sum(r["status"] == "timeout" for r in rows) / n, + clearance=float(np.mean(cl)) if cl else None, + time=float(np.mean(tt)) if tt else None, + ) + + +def phase0(out, policy, gp, beta, device, executor, cache): + rows = {"F": [], "B8": []} + support = [] + for gamma in SP.GAMMAS: + for ep in range(PH0_EP0, PH0_EP0 + 10): + for mode, key in (("full", "F"), ("b8", "B8")): + r = AF.run_certified_episode( + policy, ep, gamma, mode=mode, gp=gp, beta=beta, + device=device, executor=executor, cache=cache) + rows[key].append(dict( + episode=ep, gamma=gamma, status=r["status"], + steps=r["steps"], min_clearance=r["min_clearance"])) + if mode == "full": + support.append(dict( + episode=ep, gamma=gamma, + pool_pos=[c for c in r["pool_positive_counts"] + if c is not None])) + print(json.dumps(dict(phase0_gamma=gamma, + F=_summ([x for x in rows["F"] + if x["gamma"] == gamma]))), flush=True) + pos_counts = [c for s in support for c in s["pool_pos"]] + report = dict( + F=_summ(rows["F"]), B8=_summ(rows["B8"]), + p_pool_has_positive=float(np.mean([c > 0 for c in pos_counts])), + mean_pool_positive_multiplicity=float(np.mean(pos_counts)), + rows=rows, cache=dict(hits=cache.hits, misses=cache.misses), + ) + with open(os.path.join(out, "phase0.json"), "w") as s: + json.dump(report, s, indent=1, allow_nan=False, default=float) + print(json.dumps(dict(F=report["F"], B8=report["B8"], + p_pos=report["p_pool_has_positive"])), flush=True) + return report + + +def gather_round(policy, prev_shard, round_index, ess, device, executor, + cache, out_path): + gp, _ = AF.build_round_gp(_phi(policy), prev_shard, device=device, + ess_target=ess, round_seed=round_index) + scen = SP.expansion_scenarios(round_index) + beta, ach = _beta(policy, gp, list(scen), list(SP.GAMMAS)[:2], device, + ess, executor) + shard = AF.CertifiedQueryShard(round_index) + outcomes = [] + for s in scen: + for g in SP.GAMMAS: + r = AF.run_certified_episode( + policy, s, g, mode="b8", gp=gp, beta=beta, device=device, + executor=executor, cache=cache, shard=shard) + outcomes.append(r["status"]) + shard.save(out_path) + return shard, dict(beta=beta, ess=ach, + outcomes={k: outcomes.count(k) for k in set(outcomes)}, + Dplus=len(shard.Dplus), D=len(shard.windows)) + + +def run(args): + device = "cuda:0" + out = os.path.abspath(args.outdir) + os.makedirs(out, exist_ok=True) + policy, _ = GPS.load_sfm_policy(R1, device=device) + policy.eval() + b9 = OS.ExecutedRoundShard.load(B9_SHARD) + cache = AF.VerifierCache() + with ProcessPoolExecutor(max_workers=args.workers) as executor: + if not os.path.isfile(os.path.join(out, "phase0.json")): + gp0, _ = AF.build_round_gp(_phi(policy), b9, device=device, + ess_target=0.5, round_seed=0) + beta0, _ = _beta(policy, gp0, [PH0_EP0], list(SP.GAMMAS), + device, 0.5, executor) + phase0(out, policy, gp0, beta0, device, executor, cache) + # ---- screening ---- + arms = [] + for ess in (0.25, 0.5): + tag = str(ess).replace(".", "p") + spath = os.path.join(out, f"screen_shard_ess{tag}.pt") + if os.path.isfile(spath): + shard = AF.CertifiedQueryShard.load(spath) + ginfo = {} + else: + shard, ginfo = gather_round( + policy, b9, 2, ess, device, executor, cache, spath) + print(json.dumps(dict(screen_gather=tag, **ginfo)), + flush=True) + for obj in ("U", "G"): + for dose in DOSES: + name = f"ess{tag}_{obj}_{dose['name']}" + ck = os.path.join(out, "arms", f"{name}.pt") + if not os.path.isfile(ck): + p2, _ = GPS.load_sfm_policy(R1, device=device) + BS.configure_expansion_trainability(p2) + opt = torch.optim.Adam( + [q for q in p2.parameters() + if q.requires_grad], lr=dose["lr"]) + info = AF.certified_replay( + p2, opt, [b9, shard], objective=obj, + epochs=dose["epochs"], batch=128, + device=device, seed=AF.SEED) + os.makedirs(os.path.dirname(ck), exist_ok=True) + BX._save_checkpoint(p2, ck, dict(arm=name, + info=info)) + del p2 + torch.cuda.empty_cache() + arms.append(dict(name=name, checkpoint=ck)) + specs = [(R1, os.path.join(out, "eval_r1"))] + [ + (a["checkpoint"], os.path.join(out, f"eval_{a['name']}")) + for a in arms] + metrics = CC.evaluate_checkpoints( + specs, cache_dir=os.path.join(out, "m10_cache"), + workers_each=args.eval_workers, wave=args.eval_wave, gpu=args.gpu) + r1m = metrics[R1] + table = [dict(name="r1", m10=r1m)] + [ + dict(name=a["name"], m10=metrics[a["checkpoint"]], + checkpoint=a["checkpoint"]) for a in arms] + + def dominated(row): + m = row["m10"] + for o in table: + if o is row: + continue + q = o["m10"] + ge = (q["CR"] <= m["CR"] and q["Validity"] >= m["Validity"] + and (q["clearance"] or 0) >= (m["clearance"] or 0)) + gt = (q["CR"] < m["CR"] or q["Validity"] > m["Validity"] + or (q["clearance"] or 0) > (m["clearance"] or 0)) + if ge and gt: + return True + return False + + order = sorted( + [r for r in table if r["name"] != "r1" and not dominated(r)], + key=lambda r: (r["m10"]["CR"], -r["m10"]["Validity"], + -(r["m10"]["clearance"] or 0), -r["m10"]["SR"], + r["m10"]["timeout"], + r["m10"]["time"] or 1e9)) + promoted = order[:2] + with open(os.path.join(out, "SCREEN.json"), "w") as s: + json.dump(dict(table=table, + promoted=[p["name"] for p in promoted]), + s, indent=1, allow_nan=False, default=float) + print(json.dumps(dict(promoted=[p["name"] for p in promoted], + r1=r1m)), flush=True) + + +def main(argv=None): + ap = argparse.ArgumentParser() + ap.add_argument("--outdir", required=True) + ap.add_argument("--workers", type=int, default=28) + ap.add_argument("--eval-workers", type=int, default=12) + ap.add_argument("--eval-wave", type=int, default=5) + ap.add_argument("--gpu", type=int, default=3) + run(ap.parse_args(argv)) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/claude_afe_expand.py b/overnight_run_07_12_sfm/claude_afe_expand.py new file mode 100644 index 0000000..fa51028 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_afe_expand.py @@ -0,0 +1,117 @@ +"""Expansion stage + M50/M100 funnel for the two promoted AFE recipes.""" +from __future__ import annotations + +import argparse +from concurrent.futures import ProcessPoolExecutor +import json +import os + +import torch + +import _paths # noqa: F401 +import claude_afe_driver as AD +import claude_afe_fidelity as AF +import claude_continuation as CC +import claude_corrected_distill as CD +import grid_policy_sfm as GPS +import sfm_b1_expand as BX +import sfm_b1_offline_store as OS +import sfm_b1_store as BS + +RECIPES = (dict(name="ess0p5_U_d1", ess=0.5), dict(name="ess0p25_U_d1", ess=0.25)) +KEY = lambda m: (m["CR"], -m["Validity"], -(m["clearance"] or 0), + -m["SR"], m["timeout"], m["time"] or 1e9) + + +def run(args): + out = os.path.abspath(args.outdir) + device = "cuda:0" + cache = AF.VerifierCache() + with ProcessPoolExecutor(max_workers=args.workers) as ex: + for rec in RECIPES: + rdir = os.path.join(out, f"expand_{rec['name']}") + os.makedirs(rdir, exist_ok=True) + policy, _ = GPS.load_sfm_policy(AD.R1, device=device) + BS.configure_expansion_trainability(policy) + opt = torch.optim.Adam( + [p for p in policy.parameters() if p.requires_grad], lr=1e-5) + prev = OS.ExecutedRoundShard.load(AD.B9_SHARD) + for k in range(1, 11): + ck = os.path.join(rdir, f"round_{k:02d}.pt") + sp = os.path.join(rdir, f"shard_{k:02d}.pt") + if os.path.isfile(ck): + policy, _ = GPS.load_sfm_policy(ck, device=device) + BS.configure_expansion_trainability(policy) + opt = torch.optim.Adam([p for p in policy.parameters() + if p.requires_grad], lr=1e-5) + prev = AF.CertifiedQueryShard.load(sp) + continue + policy.eval() + shard, ginfo = AD.gather_round( + policy, prev, 1 + k, rec["ess"], device, ex, cache, sp) + info = AF.certified_replay( + policy, opt, [prev, shard], objective="U", epochs=4, + batch=128, device=device, seed=AF.SEED + k) + BX._save_checkpoint(policy, ck, dict(recipe=rec, round=k)) + print(json.dumps(dict(recipe=rec["name"], round=k, + **{x: ginfo[x] for x in ("Dplus", "D")}, + loss=info["losses"][-1])), flush=True) + prev = shard + # M10 evolution + selection + table = {} + for rec in RECIPES: + rdir = os.path.join(out, f"expand_{rec['name']}") + specs = [(AD.R1, os.path.join(out, "eval_r1"))] + [ + (os.path.join(rdir, f"round_{k:02d}.pt"), + os.path.join(rdir, f"eval_r{k}")) for k in (1, 2, 5, 10)] + met = CC.evaluate_checkpoints( + specs, cache_dir=os.path.join(out, "m10_cache"), + workers_each=args.eval_workers, wave=5, gpu=args.gpu) + rows = {("r0" if p == AD.R1 else p.split("_")[-1][:-3]): met[p] + for p, _ in specs} + table[rec["name"]] = rows + best = min( + [k for k in rows if k != "r0"], key=lambda k: KEY(rows[k])) + table[rec["name"] + "_best"] = best + with open(os.path.join(out, "EXPANSION_M10.json"), "w") as s: + json.dump(table, s, indent=1, allow_nan=False, default=float) + # M50 both, winner, M100 + finals = {} + for rec in RECIPES: + best = table[rec["name"] + "_best"] + k = int(best[1:]) if best != "r0" else 1 + finals[rec["name"]] = os.path.join( + out, f"expand_{rec['name']}", f"round_{k:02d}.pt") + m50specs = [(AD.R1, os.path.join(out, "m50_r1"))] + [ + (p, os.path.join(out, f"m50_{n}")) for n, p in finals.items()] + AD.M10_EP0, AD.M10_SEED = AD.M50_EP0, AD.M50_SEED + CC.DEV_EP0, CC.DEV_NOISE_SEED, CC.DEV_M = AD.M50_EP0, AD.M50_SEED, 50 + m50 = CC.evaluate_checkpoints( + m50specs, cache_dir=os.path.join(out, "m50_cache"), + workers_each=args.eval_workers, wave=3, gpu=args.gpu) + winner = min(finals, key=lambda n: KEY(m50[finals[n]])) + CC.DEV_EP0, CC.DEV_NOISE_SEED, CC.DEV_M = AD.M100_EP0, AD.M100_SEED, 100 + m100 = CC.evaluate_checkpoints( + [(AD.R1, os.path.join(out, "m100_r1")), + (finals[winner], os.path.join(out, "m100_winner"))], + cache_dir=os.path.join(out, "m100_cache"), + workers_each=args.eval_workers * 2, wave=2, gpu=args.gpu) + with open(os.path.join(out, "FUNNEL_FINAL.json"), "w") as s: + json.dump(dict(m50={n: m50[p] for n, p in finals.items()}, + m50_r1=m50[AD.R1], winner=winner, + winner_checkpoint=finals[winner], + m100_r1=m100[AD.R1], + m100_winner=m100[finals[winner]]), + s, indent=1, allow_nan=False, default=float) + print(json.dumps(dict(winner=winner, m100_r1=m100[AD.R1], + m100_winner=m100[finals[winner]]), + allow_nan=False, default=float), flush=True) + + +if __name__ == "__main__": + ap = argparse.ArgumentParser() + ap.add_argument("--outdir", required=True) + ap.add_argument("--workers", type=int, default=28) + ap.add_argument("--eval-workers", type=int, default=12) + ap.add_argument("--gpu", type=int, default=3) + run(ap.parse_args()) diff --git a/overnight_run_07_12_sfm/claude_afe_fidelity.py b/overnight_run_07_12_sfm/claude_afe_fidelity.py new file mode 100644 index 0000000..830d4d1 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_afe_fidelity.py @@ -0,0 +1,487 @@ +"""Exact-SOCP-faithful AFE over the historical privileged candidate pool. + +Golden rule: the exact full-H SOCP verifier is the ONLY execution/replay +authority (resolved AND y=1 AND full_h AND terminal_step=10). The privileged +SFM simulation generates and ranks proposals but never certifies; NVP +terminates fail-closed; privileged-safe/SOCP-negative is never relabeled. +""" +from __future__ import annotations + +import hashlib +import math +import os + +import numpy as np +import torch + +import _paths # noqa: F401 +import claude_corrected_distill as CD +import claude_mpc_pool as MP +import grid_feats as GF +import sfm_b1_offline_exec as OE +import sfm_b1_rbf as BR +import sfm_b1_store as BS +import sfm_hp_history as HH +import sfm_kazuki as KZ +import sfm_metrics2 as SM +import sfm_scene as SS + +SEED = 20260729 +ELL, LAM, CAP = 0.3259987743470518, 1.0e-2, 512 +B_BUDGET = 8 +S_PHI = 0.9 + + +def certified(result): + return bool( + result.get("resolved") and int(result.get("y", 0)) == 1 + and bool(result.get("full_h")) + and int(result.get("terminal_step", -1)) == 10 + ) + + +def keyed_rng(*parts): + payload = ":".join(str(p) for p in parts).encode() + return np.random.default_rng( + int.from_bytes(hashlib.sha256(payload).digest()[:8], "little"), + ) + + +class VerifierCache: + """Shared read-only-safe cache keyed by (context, U, gamma, version).""" + + def __init__(self): + self._data = {} + self.hits = 0 + self.misses = 0 + + @staticmethod + def key(context, controls, gamma): + h = hashlib.sha256() + for arr in (context["state"], context["ped_xy"], context["ped_vel"]): + h.update(np.ascontiguousarray(arr, np.float32).tobytes()) + h.update(np.ascontiguousarray(controls, np.float32).tobytes()) + h.update(f"{float(gamma):.8f}".encode()) + h.update(str(SM.verifier_manifest()).encode()) + return h.digest() + + def verify_many(self, context, plans, indices, executor): + results = {} + tasks = [] + for index in indices: + k = self.key(context, plans[index], context["gamma"]) + if k in self._data: + self.hits += 1 + results[index] = self._data[k] + else: + tasks.append((index, k)) + if tasks: + payloads = [ + (i, 0, context["state"], plans[i], context["ped_xy"], + context["ped_vel"], float(context["gamma"])) + for i, _ in tasks + ] + outs = {i: r for i, _, r in executor.map( + SM.verify_in_worker, payloads, + )} + for i, k in tasks: + self.misses += 1 + slim = {key: outs[i][key] for key in ( + "resolved", "y", "taskspace", "collision_free", + "certificate", "full_h", "terminal_step", + ) if key in outs[i]} + slim["diagnostics"] = dict( + slack=float(outs[i]["diagnostics"]["slack"]), + ) if outs[i].get("resolved") else {} + self._data[k] = slim + results[i] = slim + return results + + +def controller_scores(plans, context, cfg_gamma): + """Faithful replication of the committed step-filter score.""" + clear, inside, terminal, _, reach = KZ._simulate_sfm_plans( + _humans_placeholder(context), context["state"], np.stack(plans), 10, + ) + goal = np.asarray(SS.GOAL, np.float32) + gsw = float(cfg_gamma.step_filter_goal_score_weight) + cw = float(cfg_gamma.step_filter_clearance_weight) + nominal_u0 = np.asarray(plans[0][0], np.float32) + scores = [] + for i, plan in enumerate(plans): + goal_cost = ( + 0.04 * float(reach[i]) if reach[i] <= 10 + else float(np.linalg.norm(terminal[i, :2] - goal)) + ) + scores.append( + gsw * goal_cost + 0.015 * float(np.mean(plan * plan)) + + 0.02 * float(np.sum((plan[0] - nominal_u0) ** 2)) + - cw * min(float(clear[i]), 1.0) + ) + return np.asarray(scores, np.float64), clear, inside, reach + + +_HUMANS = {} + + +def _humans_placeholder(context): + return _HUMANS["current"] + + +def build_pool(policy, context, humans, device): + """Complete historical pool + keyed x0 + controller scores J.""" + _HUMANS["current"] = humans + pool = MP.build_codex_pool( + policy, context, humans, device=device, + seed_step=int(context["step"]), + ) + plans = pool["plans"] + cfg_gamma = KZ._gamma_controller_config( + MP.privileged_sfm_config(), float(context["gamma"]), + ).validate() + scores, clear, inside, reach = controller_scores( + list(plans), context, cfg_gamma, + ) + x0 = np.stack([ + keyed_rng(SEED, "afe_x0", context["scenario_id"], + f"{float(context['gamma']):.8f}", context["step"], i) + .standard_normal(int(policy.d)).astype(np.float32) + for i in range(len(plans)) + ]) + return dict( + plans=plans, x0=x0, J=scores, + privileged_feasible=pool["privileged_feasible"], + privileged_clearance=clear, + ) + + +@torch.no_grad() +def pool_features(policy, context, plans, x0, device): + hp10 = torch.as_tensor(context["hp10"], device=device)[None].float() + low = torch.as_tensor(context["low5"], device=device)[None].float() + hist = torch.as_tensor(context["hist"], device=device)[None].float() + ctx = policy.ctx_from(hp10, low, hist) + features = policy.phi_s_from_x0( + torch.as_tensor(np.stack(plans), device=device).float(), + ctx.repeat_interleave(len(plans), dim=0), + torch.as_tensor(x0, device=device).float(), s=S_PHI, + ) + return BR.l2_normalize(features) + + +class CertifiedQueryShard: + """Set-valued certified store: D = resolved B queries, D+ = ALL positives.""" + + def __init__(self, round_i): + self.round_i = int(round_i) + self.contexts = [] + self.windows = [] + + def add_context(self, **kw): + kw["context_id"] = len(self.contexts) + kw["round"] = self.round_i + self.contexts.append(kw) + return kw["context_id"] + + def add_query(self, context_id, controls, x0, result, *, J, sigma, + executed, source): + if int(result.get("y", -1)) == 1 and not certified(result): + raise ValueError("positive query must satisfy the golden rule") + row = dict( + window_id=len(self.windows), query_id=len(self.windows), + context_id=int(context_id), + controls=np.asarray(controls, np.float32), + x0=np.asarray(x0, np.float32), + y=int(result.get("y", 0)) if result.get("resolved") else 0, + resolved=bool(result.get("resolved")), + train_eligible=certified(result), + J=float(J), sigma=None if sigma is None else float(sigma), + executed=bool(executed), source=str(source), + ) + self.windows.append(row) + return row["window_id"] + + @property + def Dplus(self): + return [r for r in self.windows if r["train_eligible"]] + + @property + def Dminus(self): + return [r for r in self.windows if r["resolved"] and not r["y"]] + + def save(self, path): + os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True) + torch.save(dict(round=self.round_i, contexts=self.contexts, + windows=self.windows), path + ".tmp") + os.replace(path + ".tmp", path) + + @classmethod + def load(cls, path): + payload = torch.load(path, map_location="cpu", weights_only=False) + shard = cls(payload["round"]) + shard.contexts = payload["contexts"] + shard.windows = payload["windows"] + return shard + + +def run_certified_episode(policy, scenario, gamma, *, mode, gp, beta, device, + executor, cache, shard=None, T=180, reach=0.5): + """Closed-loop F (verify all) or B8 (AFE budget) certified controller.""" + environment = SS.scene_profile("double_density_velocity_ood") + humans = SS.make_humans(int(scenario), 0, environment["n_ped"], + tuple(environment["ped_speed_range"])) + state = np.zeros(4, np.float32) + history = HH.HpHistory() + controls_list = [] + status, min_clear = None, float("inf") + stats = dict(steps=0, nvp=False, pool_positive_counts=[], + b_positive_counts=[], sigma_all=[], sigma_selected=[], + oracle_rejected=0) + states, peds, pvels = [state.copy()], [], [] + for t in range(int(T)): + ped_xy, ped_vel = SS.collect_humans(humans) + clearance = float(np.linalg.norm( + ped_xy - state[:2][None], axis=1, + ).min() - SS.R_PED) + min_clear = min(min_clear, clearance) + if clearance < 0.0: + status = "collision" + break + if float(np.linalg.norm(state[:2] - SS.GOAL)) < float(reach): + status = "success" + break + obstacles = np.concatenate([ + ped_xy, np.full((len(ped_xy), 1), SS.R_PED, np.float32), + ], axis=1) + hp10 = history.append(torch.as_tensor(GF.axis_grid( + state[:2], obstacles, 0.0, R=SS.R_SENSE, sensing=SS.R_SENSE, + ))) + context = dict( + scenario_id=int(scenario), gamma=float(gamma), step=int(t), + state=state.copy(), + hp10=hp10.numpy().astype(np.float32), + low5=np.asarray(GF.low5(state, SS.GOAL, gamma), np.float32), + hist=np.asarray(GF.hist_pad( + np.asarray(controls_list[-16:]) if controls_list + else np.zeros((0, 2)), 16, + ), np.float32), + ped_xy=ped_xy.copy(), ped_vel=ped_vel.copy(), + ) + pool = build_pool(policy, context, humans, device) + n = len(pool["plans"]) + if mode == "full": + queried = list(range(n)) + sigmas = [None] * n + else: + features = pool_features( + policy, context, pool["plans"], pool["x0"], device, + ) + generator = torch.Generator(device=features.device) + generator.manual_seed(int(keyed_rng( + SEED, "acq", scenario, f"{gamma:.8f}", t, + ).integers(0, 2**62))) + selected, trace = gp.sequential_acquire( + features, min(B_BUDGET, n), beta, generator=generator, + ) + queried = list(map(int, selected)) + stats["sigma_all"].extend( + float(v) for v in + trace[0]["scores"].clamp_min(0).sqrt()[:64] + ) + stats["sigma_selected"].extend( + float(r["chosen_sigma"]) for r in trace + ) + sigmas = {q: float(r["chosen_sigma"]) + for q, r in zip(queried, trace)} + results = cache.verify_many(context, pool["plans"], queried, executor) + positives = [i for i in queried if certified(results[i])] + stats["pool_positive_counts"].append( + len(positives) if mode == "full" else None, + ) + stats["b_positive_counts"].append(len(positives)) + context_id = None + if shard is not None: + context_id = shard.add_context(**context) + for i in queried: + if results[i].get("resolved"): + shard.add_query( + context_id, pool["plans"][i], pool["x0"][i], + results[i], J=pool["J"][i], + sigma=(sigmas[i] if isinstance(sigmas, dict) + else None), + executed=False, source="pool", + ) + if not positives: + status = "nvp" + stats["nvp"] = True + break + best = min(positives, key=lambda i: (pool["J"][i], i)) + oracle_best = int(np.argmin(pool["J"])) + if oracle_best not in positives: + stats["oracle_rejected"] += 1 + if shard is not None: + for row in shard.windows[::-1]: + if row["context_id"] != context_id: + break + if np.array_equal(row["controls"], pool["plans"][best]): + row["executed"] = True + break + action = np.asarray(pool["plans"][best][0], np.float32) + controls_list.append(action) + state = state.copy() + state[:2] += SS.DT * state[2:4] + 0.5 * SS.DT ** 2 * action + state[2:4] += SS.DT * action + states.append(state.copy()) + peds.append(ped_xy) + pvels.append(ped_vel) + SS.advance_humans(humans, state) + if status is None: + status = "timeout" + stats.update( + status=status, steps=len(controls_list), + min_clearance=float(min_clear), + states=np.asarray(states, np.float32), + controls=np.asarray(controls_list, np.float32), + ped_xy=(np.asarray(peds, np.float32) if peds + else np.zeros((0, 0, 2), np.float32)), + ped_vel=(np.asarray(pvels, np.float32) if pvels + else np.zeros((0, 0, 2), np.float32)), + ) + return stats + + +def build_round_gp(phi_policy, previous_shard, *, device, ess_target, + round_seed): + """GP from previous round's certified positives; graceful gamma quota.""" + gp = BR.RBFGP(float(ELL), float(LAM)) + identities = [] + if previous_shard is not None: + by_gamma = {} + for row in previous_shard.Dplus: + context = previous_shard.contexts[int(row["context_id"])] + by_gamma.setdefault( + round(float(context["gamma"]), 8), [], + ).append(row) + quota = CAP // 7 + chosen = [] + for gamma in sorted(by_gamma): + rows = sorted( + by_gamma[gamma], + key=lambda r: (int(r["context_id"]), int(r["query_id"])), + ) + chosen.extend(rows[:quota]) + parts = [] + for start in range(0, len(chosen), 256): + values = chosen[start:start + 256] + records = [(previous_shard, row) for row in values] + grid, low, hist, controls = BS._tensor_batch(records, device) + x0 = torch.as_tensor(np.stack([ + np.asarray(row["x0"], np.float32) for row in values + ]), device=device) + parts.append(phi_policy.phi_s_from_x0( + controls, phi_policy.ctx_from(grid, low, hist), x0, s=S_PHI, + )) + if parts: + gp.set_buffer(BR.l2_normalize(torch.cat(parts))) + identities = [int(row["query_id"]) for row in chosen] + return gp, identities + + +def calibrate_tau_j(shard): + """tau_J so the within-context median ESS/|P_c| is 0.5 (bisection).""" + groups = {} + for row in shard.Dplus: + groups.setdefault(int(row["context_id"]), []).append(float(row.get("J", 0.0))) + multi = [np.asarray(v) for v in groups.values() if len(v) > 1] + if not multi: + return None, 1.0 + + def median_ess(tau): + values = [] + for J in multi: + q = np.exp(-(J - J.min()) / max(tau, 1e-9)) + q = q / q.sum() + values.append(1.0 / (np.square(q).sum() * len(q))) + return float(np.median(values)) + + lo, hi = 1e-6, 1e6 + if median_ess(hi) < 0.5: + return hi, median_ess(hi) + for _ in range(80): + mid = math.sqrt(lo * hi) + if median_ess(mid) < 0.5: + lo = mid + else: + hi = mid + return hi, median_ess(hi) + + +def objective_weights(shards, objective): + """Per-record weights over the W=2 certified union; sums to 1.""" + records = [(s, row) for s in shards for row in s.Dplus] + hmass, accounting = BS.hierarchy_mass(records) + if objective == "U": + weights = {k: float(v) for k, v in hmass.items()} + tau = None + else: + weights = {} + context_mass = {} + context_rows = {} + for shard, row in records: + key = (id(shard), int(row["context_id"])) + context_mass[key] = context_mass.get(key, 0.0) + float( + hmass[(id(shard), int(row["query_id"]))], + ) + context_rows.setdefault(key, []).append((shard, row)) + tau = None + taus = [calibrate_tau_j(s)[0] for s in shards] + taus = [t for t in taus if t is not None] + tau = float(np.median(taus)) if taus else 1.0 + for key, rows in context_rows.items(): + J = np.asarray([float(r.get("J", 0.0)) for _, r in rows]) + q = np.exp(-(J - J.min()) / max(tau, 1e-9)) + q = q / q.sum() + for (shard, row), qi in zip(rows, q): + weights[(id(shard), int(row["query_id"]))] = ( + context_mass[key] * float(qi) + ) + residual = abs(sum(weights.values()) - 1.0) + return records, weights, dict( + objective=objective, tau_J=tau, mass_residual=residual, + gamma=accounting["gamma"], + ) + + +def certified_replay(policy, optimizer, shards, *, objective, epochs, batch, + device, seed): + """Whole-dataset accumulated replay; one Adam step per epoch.""" + records, weights, info = objective_weights(shards, objective) + if info["mass_residual"] > 1e-6: + raise RuntimeError(f"hierarchy-mass residual {info['mass_residual']}") + encoder = BS.module_sha256(policy.enc_grid) + policy.train() + losses = [] + n = len(records) + for epoch in range(int(epochs)): + optimizer.zero_grad(set_to_none=True) + total = 0.0 + for start in range(0, n, int(batch)): + values = records[start:start + int(batch)] + grid, low, hist, controls = BS._tensor_batch(values, device) + ctx = policy.ctx_from(grid, low, hist) + w = torch.as_tensor([ + len(values) * weights[(id(h), int(r["query_id"]))] + for h, r in values + ], dtype=controls.dtype, device=device) + torch.manual_seed(int(seed) + epoch * 1_000_003 + start) + loss = policy.cfm_loss(controls, ctx, weights=w) + loss.backward() + total += float(loss.detach()) + optimizer.step() + losses.append(total) + policy.eval() + if BS.module_sha256(policy.enc_grid) != encoder: + raise RuntimeError("visual encoder changed") + info.update(adam_steps=int(epochs), losses=losses, records=n, + coverage=n) + return info diff --git a/overnight_run_07_12_sfm/claude_apply_unverified_mpc_teacher.py b/overnight_run_07_12_sfm/claude_apply_unverified_mpc_teacher.py new file mode 100644 index 0000000..5057ab9 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_apply_unverified_mpc_teacher.py @@ -0,0 +1,320 @@ +"""Apply one isolated privileged-MPC teacher block after an ordinary SFE update. + +The input checkpoint is treated as ``theta_(n+1/2)``: ordinary executed +``D+/D-`` replay has already happened. This command harvests a separate +control-bounded ``D_MPC`` from the matching round shard, applies a dedicated +CFM block, and writes ``theta_(n+1)``. It never edits the round shard and never +routes teacher rows into the verifier, GP, acquisition, or validity metrics. +""" +from __future__ import annotations + +import argparse +from concurrent.futures import ProcessPoolExecutor +import json +import os + +import numpy as np +import torch + +import _paths # noqa: F401 +import claude_offline_aug as AUG +import claude_unverified_mpc_teacher as T +import grid_policy_sfm as GPS +import sfm_b1_expand as BX +import sfm_b1_offline_exec as OE +import sfm_b1_offline_store as OS +import sfm_b1_store as BS +import sfm_scene as SS + + +STATUS = "SFM_UNVERIFIED_MPC_TEACHER_BLOCK_COMPLETE" + + +def _write_json(path, payload): + os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True) + temporary = os.fspath(path) + ".tmp" + with open(temporary, "w") as stream: + json.dump(payload, stream, indent=2, sort_keys=True, allow_nan=False) + os.replace(temporary, path) + + +def _hard_windows(shard): + _, population_b, stats = AUG.tag_populations(shard) + rows = {int(row["window_id"]): row for row in population_b} + for row in shard.windows: + if bool(row.get("nvp_context")): + rows[int(row["window_id"])] = row + return [rows[key] for key in sorted(rows)], stats + + +def _probe_objective(policy, shard, records, *, teacher, device, seed): + if not records: + return None, {} + if teacher: + values = list(records)[:128] + hp10, low, hist, controls = T._tensor_batch(shard, values, device) + else: + values = list(records)[:128] + contexts = [shard.contexts[int(row["context_id"])] for row in values] + hp10 = torch.as_tensor( + np.stack([row["hp10"] for row in contexts]), device=device, + ).float() + low = torch.as_tensor( + np.stack([row["low5"] for row in contexts]), device=device, + ).float() + hist = torch.as_tensor( + np.stack([row["hist"] for row in contexts]), device=device, + ).float() + controls = torch.as_tensor( + np.stack([row["controls"] for row in values]), device=device, + ).float() + torch.manual_seed(int(seed)) + context = policy.ctx_from(hp10, low, hist) + loss = policy.cfm_loss(controls, context) + parameters = [ + parameter for parameter in policy.parameters() + if parameter.requires_grad + ] + gradients = torch.autograd.grad( + loss, parameters, allow_unused=True, retain_graph=False, + ) + snapshot = { + name: gradient.detach().cpu() + for (name, parameter), gradient in zip( + ( + (name, parameter) + for name, parameter in policy.named_parameters() + if parameter.requires_grad + ), + gradients, + ) + if gradient is not None + } + return float(loss.detach()), snapshot + + +def _gradient_cosine(left, right): + common = sorted(set(left) & set(right)) + if not common: + return None + numerator = sum( + float((left[name].double() * right[name].double()).sum()) + for name in common + ) + left_norm = sum( + float(left[name].double().square().sum()) for name in common + ) ** 0.5 + right_norm = sum( + float(right[name].double().square().sum()) for name in common + ) ** 0.5 + if left_norm == 0.0 or right_norm == 0.0: + return None + return float(numerator / (left_norm * right_norm)) + + +def _conflict_probe(policy, shard, teachers, *, device, seed): + was_training = policy.training + # cuDNN GRU backward is unavailable in eval mode. Keep the probe in + # seeded train mode; the same seed and records are reused before/after. + policy.train() + positive_loss, positive_gradient = _probe_objective( + policy, shard, shard.Dplus, teacher=False, + device=device, seed=seed, + ) + negative_loss, negative_gradient = _probe_objective( + policy, shard, shard.Dminus, teacher=False, + device=device, seed=seed, + ) + teacher_loss, teacher_gradient = _probe_objective( + policy, shard, teachers, teacher=True, + device=device, seed=seed, + ) + policy.train(was_training) + return { + "ordinary_positive_loss": positive_loss, + "ordinary_negative_loss": negative_loss, + "teacher_loss": teacher_loss, + "cos_teacher_positive": _gradient_cosine( + teacher_gradient, positive_gradient, + ), + "cos_teacher_negative": _gradient_cosine( + teacher_gradient, negative_gradient, + ), + "probe_seed": int(seed), + "ordinary_positive_records": min(len(shard.Dplus), 128), + "ordinary_negative_records": min(len(shard.Dminus), 128), + "teacher_records": min(len(teachers), 128), + } + + +def run(args): + checkpoint = os.path.abspath(args.checkpoint) + round_shard_path = os.path.abspath(args.round_shard) + output_dir = os.path.abspath(args.output_dir) + if os.path.exists(output_dir): + raise FileExistsError(output_dir) + checkpoint_sha = OS.sha256_file(checkpoint) + if checkpoint_sha != args.expected_checkpoint_sha256: + raise ValueError( + "checkpoint SHA mismatch: " + f"expected {args.expected_checkpoint_sha256}, got {checkpoint_sha}" + ) + if args.expected_checkpoint_sha256 == T.R1_CHECKPOINT_SHA256: + parent_contract = "immutable accepted r1" + else: + parent_contract = "caller-pinned post-ordinary checkpoint" + shard = OS.ExecutedRoundShard.load(round_shard_path) + policy, _ = GPS.load_sfm_policy(checkpoint, device=args.device) + BS.configure_expansion_trainability(policy) + encoder_sha_before = BS.module_sha256(policy.enc_grid) + hard, population_stats = _hard_windows(shard) + environment = SS.scene_profile(args.scene_profile) + os.makedirs(output_dir) + if getattr(args, "reuse_buffer", None): + # Additive, default-off: reuse a previously harvested D_MPC buffer. + # Valid ONLY when it authenticates against the identical round shard; + # the harvest is deterministic given (checkpoint, shard), so this is + # byte-equivalent to re-harvesting and is unit-tested as such. + reuse_path = os.path.abspath(args.reuse_buffer) + payload = torch.load(reuse_path, map_location="cpu", + weights_only=False) + if payload.get("status") != T.BUFFER_STATUS: + raise ValueError("reused buffer is not a complete teacher buffer") + if payload.get("round_shard_sha256") != OS.sha256_file( + round_shard_path, + ): + raise ValueError( + "reused buffer does not authenticate against this round shard" + ) + records = list(payload["records"]) + # The audit is kept byte-identical so the re-saved buffer hashes + # exactly like the source buffer; reuse provenance goes into the + # COMPLETE.json report instead. + harvest = dict(payload["audit"]) + else: + with ProcessPoolExecutor( + max_workers=args.verifier_workers, + ) as executor: + records, harvest = T.harvest_round( + policy, + shard, + hard, + device=args.device, + environment=environment, + max_contexts=args.max_contexts, + executor=executor if args.audit_socp else None, + audit_socp=args.audit_socp, + ) + buffer_path = os.path.join(output_dir, "D_MPC.pt") + T.save_buffer( + buffer_path, + round_shard_path, + shard, + records, + harvest, + ) + probe_before = _conflict_probe( + policy, + shard, + records, + device=args.device, + seed=args.seed, + ) + optimizer = torch.optim.Adam( + [parameter for parameter in policy.parameters() + if parameter.requires_grad], + lr=args.teacher_lr, + ) + update = T.distill_block( + policy, + optimizer, + shard, + records, + epochs=args.teacher_epochs, + batch=args.batch, + seed=args.seed, + ) + if BS.module_sha256(policy.enc_grid) != encoder_sha_before: + raise RuntimeError("visual encoder changed during teacher block") + probe_after = _conflict_probe( + policy, + shard, + records, + device=args.device, + seed=args.seed, + ) + output_checkpoint = os.path.join(output_dir, "post_teacher.pt") + BX._save_checkpoint( + policy, + output_checkpoint, + { + "phase": "post_ordinary_then_unverified_privileged_mpc_teacher", + "source_checkpoint": checkpoint, + "source_checkpoint_sha256": checkpoint_sha, + "round_shard": round_shard_path, + "round_shard_sha256": OS.sha256_file(round_shard_path), + "teacher_buffer": buffer_path, + "teacher_lr": float(args.teacher_lr), + "teacher_epochs": int(args.teacher_epochs), + "teacher_seed": int(args.seed), + }, + ) + report = { + "status": STATUS, + "parent_contract": parent_contract, + "source_checkpoint": checkpoint, + "source_checkpoint_sha256": checkpoint_sha, + "round": int(shard.round_i), + "round_shard": round_shard_path, + "round_shard_sha256": OS.sha256_file(round_shard_path), + "teacher_buffer": buffer_path, + "teacher_buffer_sha256": OS.sha256_file(buffer_path), + "post_teacher_checkpoint": output_checkpoint, + "post_teacher_checkpoint_sha256": OS.sha256_file(output_checkpoint), + "ordinary_replay_precedes_this_command": True, + "teacher_is_not_a_safety_label": True, + "teacher_used_at_evaluation": False, + "reused_buffer_from": getattr(args, "reuse_buffer", None), + "population_stats": population_stats, + "harvest": harvest, + "update": update, + "conflict_probe_before": probe_before, + "conflict_probe_after": probe_after, + } + _write_json(os.path.join(output_dir, "COMPLETE.json"), report) + return report + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--checkpoint", required=True) + parser.add_argument( + "--expected-checkpoint-sha256", + default=T.R1_CHECKPOINT_SHA256, + ) + parser.add_argument("--round-shard", required=True) + parser.add_argument("--output-dir", required=True) + parser.add_argument("--scene-profile", default=OE.SCENE_PROFILE) + parser.add_argument("--teacher-lr", type=float, required=True) + parser.add_argument("--teacher-epochs", type=int, required=True) + parser.add_argument("--batch", type=int, default=128) + parser.add_argument("--seed", type=int, default=20260728) + parser.add_argument("--max-contexts", type=int, default=400) + parser.add_argument("--audit-socp", action="store_true") + parser.add_argument( + "--reuse-buffer", default=None, + help="path to a D_MPC.pt from an identical (checkpoint, shard) run; " + "skips the deterministic harvest after authenticating the shard", + ) + parser.add_argument("--verifier-workers", type=int, default=32) + parser.add_argument("--device", default="cuda") + args = parser.parse_args(argv) + if not 0.0 < float(args.teacher_lr) <= 1.0e-3: + parser.error("--teacher-lr must be in (0,1e-3]") + if not 0 <= int(args.teacher_epochs) <= 32: + parser.error("--teacher-epochs must be in [0,32]") + run(args) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/claude_confirm_analysis.py b/overnight_run_07_12_sfm/claude_confirm_analysis.py new file mode 100644 index 0000000..3cdbff6 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_confirm_analysis.py @@ -0,0 +1,160 @@ +"""Paired-difference confirmation analysis for the Stage-E M100 comparison. + +Inputs: the raw-evaluator metrics JSON (r0 + selected checkpoint on one CRN +bank) and the locked-Kazuki metrics JSON on the same episode bank. All three +methods share the same scenario ids per gamma (CRN pedestrian banks); r0 and +the selected checkpoint additionally share the same latent bank. + +For each pair (selected - r0, selected - Kazuki, Kazuki - r0) and each of the +four study metrics we report the mean difference with a 95% scenario-cluster +bootstrap interval: episodes are grouped by scenario id (keeping all seven +paired gamma rows together) and clusters are resampled with replacement. +Collision and Validity use all episodes; clearance and time-to-goal are +success-conditioned, so each bootstrap draw recomputes the per-method mean +over its successful episodes inside the resampled clusters (a paired +difference of success-conditioned means, not a per-episode paired delta). +""" +from __future__ import annotations + +import argparse +import json + +import numpy as np + + +METRICS = ("CR", "validity", "successful_clearance", "time_to_goal") + + +def _rows(path, record_index=None, expect_label=None): + with open(path) as stream: + payload = json.load(stream) + if "records" in payload: + record = payload["records"][record_index] + if expect_label is not None and record["label"] != expect_label: + raise ValueError( + f"{path}: expected label {expect_label}, got {record['label']}" + ) + return payload, record["cell"]["rows"] + return payload, payload["rows"] + + +def _by_scenario(rows): + grouped = {} + for row in rows: + grouped.setdefault(int(row["episode"]), []).append(row) + return grouped + + +def _metric_values(rows, metric): + if metric == "CR": + return [float(bool(row["collision"])) for row in rows] + if metric == "validity": + return [float(row["validity"]) for row in rows] + key = ( + "successful_clearance" if metric == "successful_clearance" + else "time_to_goal" + ) + return [ + float(row[key]) for row in rows if row[key] is not None + ] + + +def _cluster_mean(grouped, scenarios, metric): + values = [] + for scenario in scenarios: + values.extend(_metric_values(grouped[scenario], metric)) + return float(np.mean(values)) if values else float("nan") + + +def paired_difference(rows_a, rows_b, *, seed, draws=10_000): + grouped_a, grouped_b = _by_scenario(rows_a), _by_scenario(rows_b) + scenarios = sorted(set(grouped_a) & set(grouped_b)) + if set(grouped_a) != set(grouped_b): + raise ValueError("methods do not share the scenario bank") + generator = np.random.default_rng(seed) + out = {} + for metric in METRICS: + point = ( + _cluster_mean(grouped_a, scenarios, metric) + - _cluster_mean(grouped_b, scenarios, metric) + ) + samples = [] + for _ in range(draws): + resample = generator.choice(scenarios, size=len(scenarios)) + samples.append( + _cluster_mean(grouped_a, resample, metric) + - _cluster_mean(grouped_b, resample, metric) + ) + finite = [s for s in samples if np.isfinite(s)] + low, high = np.quantile(finite, [0.025, 0.975]) + out[metric] = dict( + difference=point, ci95=[float(low), float(high)], + draws=len(finite), + ) + return out + + +def summarize_method(rows): + n = len(rows) + values = {m: _metric_values(rows, m) for m in METRICS} + return dict( + n=n, + SR=float(np.mean([bool(r["success"]) for r in rows])), + CR=float(np.mean(values["CR"])), + timeout=float(np.mean([bool(r["timeout"]) for r in rows])), + Validity=float(np.mean(values["validity"])), + successful_clearance=float(np.mean(values["successful_clearance"])), + successful_time_to_goal=float(np.mean(values["time_to_goal"])), + successes=int(sum(bool(r["success"]) for r in rows)), + ) + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--raw-metrics", required=True, + help="raw M100 metrics json containing r0 + selected") + parser.add_argument("--kazuki-metrics", required=True) + parser.add_argument("--selected-label", required=True) + parser.add_argument("--seed", type=int, default=20260731) + parser.add_argument("--out", required=True) + args = parser.parse_args(argv) + + raw_payload, r0_rows = _rows(args.raw_metrics, 0, "r0") + _, selected_rows = _rows(args.raw_metrics, 1, args.selected_label) + kazuki_payload, kazuki_rows = _rows(args.kazuki_metrics) + + result = dict( + status="CLAUDE_M100_CONFIRMATION_ANALYSIS", + bank=raw_payload.get("bank"), + kazuki_config=kazuki_payload.get("kazuki_config"), + methods=dict( + r0=summarize_method(r0_rows), + selected=summarize_method(selected_rows), + kazuki=summarize_method(kazuki_rows), + ), + paired_differences=dict( + selected_minus_r0=paired_difference( + selected_rows, r0_rows, seed=args.seed, + ), + selected_minus_kazuki=paired_difference( + selected_rows, kazuki_rows, seed=args.seed + 1, + ), + kazuki_minus_r0=paired_difference( + kazuki_rows, r0_rows, seed=args.seed + 2, + ), + ), + semantics=( + "scenario-cluster bootstrap (10k draws) keeping the seven paired " + "gamma rows per scenario together; clearance/time are " + "success-conditioned means recomputed inside each draw" + ), + ) + with open(args.out, "w") as stream: + json.dump(result, stream, indent=2, allow_nan=False) + print(json.dumps(result["methods"], indent=1)) + print(json.dumps(result["paired_differences"]["selected_minus_r0"], + indent=1)) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/claude_continuation.py b/overnight_run_07_12_sfm/claude_continuation.py new file mode 100644 index 0000000..a2490bf --- /dev/null +++ b/overnight_run_07_12_sfm/claude_continuation.py @@ -0,0 +1,679 @@ +"""Stable continuation from the confirmed r1 checkpoint (pre-registered). + +Development mode: per continuation round, gather ONE shard from the currently +accepted checkpoint (unchanged margin-selector B1 protocol, next expansion +scenario block), build the round's exact-certified recovery positives, then +fork 16 candidates (E x lr x anchor_mass) from the accepted checkpoint. All +candidates share the identical shard, recovery records, replay seed, and +deterministic batch ordering. Positive replay mass composition (declared): + + anchor_mass * r1-self-anchor + 0.05 * recovery + + (0.95 - anchor_mass) * new-shard D+ + +with the standard hierarchy mass inside each population and the unchanged +alpha=0.01 signed-gradient negative scheme on the new shard's D-. Every +candidate is evaluated on the fixed raw M10/gamma development bank; the +pre-registered hard admissibility gate versus the immutable r1 baseline +applies, the lexicographic rule selects the accepted checkpoint, and the +procedure stops when nothing is admissible. + +Confirmation mode: identical loop with one fixed (E, lr, anchor_mass) combo, +restarted from immutable r1, no round-dependent tuning. +""" +from __future__ import annotations + +import argparse +from concurrent.futures import ProcessPoolExecutor +import copy +from dataclasses import dataclass +import itertools +import json +import math +import os +import subprocess +import sys +import time + +import numpy as np +import torch + +import _paths # noqa: F401 +import claude_offline_aug as AUG +import grid_policy_sfm as GPS +import sfm_b1_eval as BE +import sfm_b1_expand as BX +import sfm_b1_offline_exec as OE +import sfm_b1_offline_store as OS +import sfm_b1_r2_alpha_replay as R2 +import sfm_b1_store as BS +import sfm_metrics2 as SM +import sfm_protocol as SP +import sfm_scene as SS + +E_GRID = (1, 4, 10, 25) +LR_GRID = (1e-5, 3e-5) +ANCHOR_GRID = (0.25, 0.50) +RECOVERY_SHARE = 0.05 +DEV_EP0, DEV_NOISE_SEED, DEV_M = 390_000, 20_260_740, 10 +BATCH = 128 +ALPHA = 0.01 +SEED = 20260724 +SR_TOL = 0.05 +TIMEOUT_TOL = 0.05 +B9_ELL = 0.3259987743470518 +HERE = os.path.dirname(os.path.abspath(__file__)) + + +@dataclass(frozen=True) +class ContOpts: + K: int = 16 + B: int = 4 + T: int = 180 + H: int = 10 + nfe: int = 8 + temp: float = 1.0 + phi_s: float = 0.9 + gp_lam: float = OE.GP_LAMBDA + ess_target: float = 0.5 + seed: int = SEED + selector: str = "margin" + smoke: bool = False + scene_profile: str = OE.SCENE_PROFILE + + +class AnchorShard: + """Duck-typed frozen r1 self-anchor buffer for the mass machinery.""" + + def __init__(self, payload): + self.round_i = 0 + self.contexts, self.windows = [], [] + for index, record in enumerate(payload["records"]): + context = record["context"] + self.contexts.append(dict( + context_id=index, round=0, + scenario_id=int(record["episode"]), + gamma=float(record["gamma"]), step=int(record["step"]), + state=np.asarray(context["state"], np.float32), + hp10=np.asarray(context["hp10"], np.float32), + low5=np.asarray(context["low5"], np.float32), + hist=np.asarray(context["hist"], np.float32), + ped_xy=np.asarray(context["ped_xy"], np.float32), + ped_vel=np.asarray(context["ped_vel"], np.float32), + )) + self.windows.append(dict( + window_id=index, query_id=index, context_id=index, + controls=np.asarray(record["controls"], np.float32), y=1, + )) + + +def _interleave(streams): + """Deterministic proportional interleave of any number of streams.""" + streams = [list(stream) for stream in streams if stream] + cursors = [0] * len(streams) + merged = [] + while any(c < len(s) for c, s in zip(cursors, streams)): + progress = [ + (cursors[i] + 1) / len(streams[i]) if cursors[i] < len(streams[i]) + else float("inf") + for i in range(len(streams)) + ] + pick = int(np.argmin(progress)) + merged.append(streams[pick][cursors[pick]]) + cursors[pick] += 1 + return merged + + +def build_positive_populations(shard, recovery_records, anchor_shard, + anchor_mass): + recovery_view = AUG.ShardView(shard, recovery_records) \ + if recovery_records else None + populations = [ + ("new", shard, [(shard, row) for row in shard.Dplus], + 0.95 - float(anchor_mass)), + ("anchor", anchor_shard, + [(anchor_shard, row) for row in anchor_shard.windows], + float(anchor_mass)), + ] + if recovery_view is not None: + populations.insert(1, ( + "recovery", recovery_view, + [(recovery_view, row) for row in recovery_view.windows], + RECOVERY_SHARE, + )) + else: + populations[0] = ( + "new", shard, populations[0][2], + 0.95 - float(anchor_mass) + RECOVERY_SHARE, + ) + return populations + + +def build_weights(populations, negatives): + total_positive = sum(len(records) for _, _, records, _ in populations) + weights = {} + membership = {} + for name, _, records, share in populations: + mass, _ = BS.hierarchy_mass(records) + for holder, row in records: + key = (id(holder), int(row["query_id"])) + weights[key] = ( + float(share) * float(mass[key]) * float(total_positive) + ) + membership[key] = name + negative_mass, _ = BS.hierarchy_mass(negatives) if negatives else ({}, {}) + for holder, row in negatives: + key = (id(holder), int(row["query_id"])) + weights[key] = float(negative_mass[key]) * float(len(negatives)) + membership[key] = "negative" + return weights, membership, total_positive + + +def epoch_batches(populations, negatives, *, batch, seed, epoch): + streams = [ + BS.hierarchical_order(records, int(seed) + epoch * 1009 + 7 * index) + for index, (_, _, records, _) in enumerate(populations) + ] + positive_stream = _interleave(streams) + negative_stream = BS.hierarchical_order( + negatives, int(seed) + epoch * 1009 + 997, + ) if negatives else [] + total = len(positive_stream) + len(negative_stream) + batch_count = math.ceil(total / int(batch)) if total else 0 + if positive_stream and len(positive_stream) < batch_count: + raise RuntimeError("cannot seed every minibatch with a positive") + batches = [[positive_stream[i]] for i in range(batch_count)] + remaining = _interleave([positive_stream[batch_count:], negative_stream]) + index = 0 + for record in remaining: + while len(batches[index]) >= int(batch): + index = (index + 1) % batch_count + batches[index].append(record) + index = (index + 1) % batch_count + return batches + + +def _weighted_cfm(policy, records, weights, device, generator_seed): + grid, low, hist, controls = BS._tensor_batch(records, device) + context = policy.ctx_from(grid, low, hist) + values = torch.as_tensor([ + weights[(id(holder), int(row["query_id"]))] + for holder, row in records + ], dtype=controls.dtype, device=device) + torch.manual_seed(int(generator_seed)) + return policy.cfm_loss(controls, context, weights=values) + + +def continuation_replay(policy, optimizer, populations, negatives, *, + epochs, batch, seed, device): + weights, membership, total_positive = build_weights( + populations, negatives, + ) + policy.train() + encoder_before = BS.module_sha256(policy.enc_grid) + module_before = R2._module_snapshot(policy) + per_pop_losses = {name: [] for name, *_ in populations} + per_pop_losses["negative"] = [] + gradient_norms = [] + steps = 0 + for epoch in range(int(epochs)): + batches = epoch_batches( + populations, negatives, batch=batch, seed=seed, epoch=epoch, + ) + for batch_index, values in enumerate(batches): + positive = [ + r for r in values + if membership[(id(r[0]), int(r[1]["query_id"]))] != "negative" + ] + negative = [ + r for r in values + if membership[(id(r[0]), int(r[1]["query_id"]))] == "negative" + ] + if not positive: + continue + step_seed = int(seed) + epoch * 100_003 + batch_index + optimizer.zero_grad(set_to_none=True) + positive_loss = _weighted_cfm( + policy, positive, weights, device, step_seed, + ) + if not bool(torch.isfinite(positive_loss)): + raise FloatingPointError("non-finite positive loss") + positive_loss.backward() + positive_gradient = BS._gradient_snapshot(policy) + positive_norm = BS._gradient_norm(positive_gradient) + rho = 0.0 + negative_gradient = {} + if ALPHA > 0.0 and negative: + optimizer.zero_grad(set_to_none=True) + negative_loss = _weighted_cfm( + policy, negative, weights, device, step_seed + 51, + ) + negative_loss.backward() + negative_gradient = BS._gradient_snapshot(policy) + negative_norm = BS._gradient_norm(negative_gradient) + rho = ALPHA * positive_norm / (negative_norm + 1e-12) + per_pop_losses["negative"].append(float(negative_loss)) + for name, parameter in policy.named_parameters(): + if not parameter.requires_grad: + continue + pos = positive_gradient.get(name) + neg = negative_gradient.get(name) + if pos is None and neg is None: + parameter.grad = None + elif pos is None: + parameter.grad = -rho * neg + elif neg is None: + parameter.grad = pos + else: + parameter.grad = pos - rho * neg + optimizer.step() + gradient_norms.append(float(positive_norm)) + steps += 1 + with torch.no_grad(): + for name, _, records, _ in populations: + subset = [ + r for r in positive + if membership[(id(r[0]), int(r[1]["query_id"]))] + == name + ] + if subset: + per_pop_losses[name].append(float(_weighted_cfm( + policy, subset, weights, device, step_seed, + ))) + policy.eval() + if BS.module_sha256(policy.enc_grid) != encoder_before: + raise RuntimeError("visual encoder changed during continuation replay") + drift = R2._module_relative_drift( + module_before, R2._module_snapshot(policy), + ) + def _summary(values): + return None if not values else dict( + first=values[0], last=values[-1], mean=float(np.mean(values)), + ) + return dict( + steps=steps, epochs=int(epochs), + unique_positive=total_positive, + unique_negative=len(negatives), + gradient_norm=_summary(gradient_norms), + losses={k: _summary(v) for k, v in per_pop_losses.items()}, + parameter_drift=drift, + ) + + +@torch.no_grad() +def probe_diagnostics(policy, reference_policy, anchor_shard, device): + """Route diversity + representation drift on fixed probe contexts.""" + by_gamma = {} + for context in anchor_shard.contexts: + by_gamma.setdefault(round(context["gamma"], 8), []).append(context) + probes = [] + for gamma in sorted(by_gamma): + probes.extend(by_gamma[gamma][:5]) + probes = probes[:35] + modes = {"yield": 0, "left": 0, "right": 0} + spreads = [] + cosines = [] + for probe in probes: + hp10 = torch.as_tensor(probe["hp10"], device=device)[None].float() + low = torch.as_tensor(probe["low5"], device=device)[None].float() + hist = torch.as_tensor(probe["hist"], device=device)[None].float() + ctx = policy.ctx_from(hp10, low, hist) + generator = np.random.default_rng(OE._keyed_seed( + SEED, 0, int(probe["scenario_id"]), f"{probe['gamma']:.8f}", + int(probe["step"]), "route_probe", + )) + x0 = generator.standard_normal((16, int(policy.d)), dtype=np.float32) + windows = BE.integrate_latents( + policy, torch.as_tensor(x0, device=device), + ctx.repeat_interleave(16, dim=0), nfe=8, + ).reshape(16, 10, 2) + windows_np = windows.cpu().numpy() + prediction = SM.predict_pedestrians( + probe["ped_xy"], probe["ped_vel"], H=10, + ) + for k in range(16): + segment = SM.rollout_positions(probe["state"], windows_np[k]) + modes[BE.classify_candidate(segment, prediction)] += 1 + spreads.append(float(np.mean(np.linalg.norm( + windows_np[:, None] - windows_np[None, :], axis=(2, 3), + )))) + features = policy.phi_s_from_x0( + windows.reshape(16, 10, 2), ctx.repeat_interleave(16, dim=0), + torch.as_tensor(x0, device=device), s=0.9, + ) + if reference_policy is not None: + ref_ctx = reference_policy.ctx_from(hp10, low, hist) + ref_windows = BE.integrate_latents( + reference_policy, torch.as_tensor(x0, device=device), + ref_ctx.repeat_interleave(16, dim=0), nfe=8, + ).reshape(16, 10, 2) + ref_features = reference_policy.phi_s_from_x0( + ref_windows, ref_ctx.repeat_interleave(16, dim=0), + torch.as_tensor(x0, device=device), s=0.9, + ) + cosine = torch.nn.functional.cosine_similarity( + features, ref_features, dim=1, + ).mean() + cosines.append(float(cosine)) + total_modes = max(sum(modes.values()), 1) + probabilities = [v / total_modes for v in modes.values() if v > 0] + entropy = -sum(p * math.log(p) for p in probabilities) + return dict( + probe_contexts=len(probes), + route_counts=modes, + route_entropy=float(entropy), + mean_pairwise_control_spread=float(np.mean(spreads)), + representation_cosine_vs_accepted=( + None if not cosines else float(np.mean(cosines)) + ), + ) + + +def m10_pooled(path): + with open(path) as stream: + payload = json.load(stream) + record = payload["records"][0] + cell = record["cell"]["summary"] + p = cell["pooled"] + def _cell(gamma): + c = cell["per_gamma"][gamma] + return dict( + clearance=c["successful_clearance"]["mean"], + time=c["successful_time_to_goal"]["mean"], + ) + return dict( + SR=float(p["SR"]), CR=float(p["CR"]), timeout=float(p["timeout"]), + Validity=float(p["Validity"]["mean"]), + clearance=p["successful_clearance"]["mean"], + time=p["successful_time_to_goal"]["mean"], + g01=_cell("0.1"), g10=_cell("1.0"), + ) + + +def admissible(candidate, r1): + checks = dict( + SR=candidate["SR"] >= r1["SR"] - SR_TOL, + timeout=candidate["timeout"] <= r1["timeout"] + TIMEOUT_TOL, + CR=candidate["CR"] <= r1["CR"], + Validity=candidate["Validity"] >= r1["Validity"], + clearance=( + candidate["clearance"] is not None + and r1["clearance"] is not None + and candidate["clearance"] >= r1["clearance"] + ), + gamma_trend=( + candidate["g01"]["clearance"] is not None + and candidate["g10"]["clearance"] is not None + and candidate["g01"]["time"] is not None + and candidate["g10"]["time"] is not None + and candidate["g01"]["clearance"] >= candidate["g10"]["clearance"] + and candidate["g01"]["time"] >= candidate["g10"]["time"] + ), + ) + return all(checks.values()), checks + + +def lex_key(candidate): + return ( + candidate["CR"], -candidate["Validity"], + -(candidate["clearance"] if candidate["clearance"] is not None + else -1.0), + candidate["time"] if candidate["time"] is not None else 1e9, + ) + + +def gather_round(policy, previous_shard, round_index, opts, device, + executor): + policy.eval() + phi_policy = copy.deepcopy(policy).eval() + for parameter in phi_policy.parameters(): + parameter.requires_grad_(False) + environment = SS.scene_profile(opts.scene_profile) + replicas = [ + BX.Replica(s, g, n_ped=environment["n_ped"], + ped_speed_range=tuple(environment["ped_speed_range"])) + for s in SP.expansion_scenarios(round_index, smoke=opts.smoke) + for g in SP.GAMMAS + ] + gp, _, _ = OE.gp_from_previous( + phi_policy, previous_shard, round_i=round_index, ell=B9_ELL, + cap=OE.CAP, lam=opts.gp_lam, phi_s=opts.phi_s, device=device, + seed=opts.seed + round_index * 101, + ) + beta, ess = OE._calibrate_beta( + phi_policy, gp, replicas, opts, device, round_i=round_index, + ) + shard = OS.ExecutedRoundShard(round_index) + gather = OE.gather_offline_round( + policy, phi_policy, gp, beta, replicas, opts, shard, device, + executor, round_i=round_index, + ) + _, pop_b, pop_stats = AUG.tag_populations(shard) + recovery, recovery_audit = AUG.build_recovery_records( + shard, pop_b, executor, family="v1", + ) + return shard, recovery, dict( + beta=float(beta), ess=float(ess), + outcomes={s: sum(o["status"] == s for o in gather["outcomes"]) + for s in ("success", "collision", "timeout")}, + counts=dict(gather["counts"]), + populations=pop_stats, + recovery_kept=recovery_audit["certified_kept"], + ) + + +def evaluate_checkpoints(specs, *, cache_dir, workers_each, wave, gpu): + """specs: list of (checkpoint, outdir). Runs waves of parallel evals.""" + results = {} + for start in range(0, len(specs), int(wave)): + processes = [] + for checkpoint, outdir in specs[start:start + int(wave)]: + if os.path.isfile(os.path.join( + outdir, f"raw_m{DEV_M}_offline_metrics.json", + )): + continue + env = dict(os.environ) + env.update( + CUDA_DEVICE_ORDER="PCI_BUS_ID", CUDA_VISIBLE_DEVICES=str(gpu), + PYTHONPATH=HERE, + ) + processes.append(subprocess.Popen( + [sys.executable, + os.path.join(HERE, "sfm_b1_offline_eval.py"), + "--checkpoints", checkpoint, "--labels", "r1", + "--scene-profile", "double_density_velocity_ood", + "--ep0", str(DEV_EP0), "--noise-seed", str(DEV_NOISE_SEED), + "--m-per-gamma", str(DEV_M), "--device", "cuda:0", + "--workers", str(int(workers_each)), + "--cache-dir", cache_dir, "--output-dir", outdir], + env=env, stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, cwd=HERE, + )) + for process in processes: + if process.wait() != 0: + raise RuntimeError("candidate M10 evaluation failed") + for checkpoint, outdir in specs: + results[checkpoint] = m10_pooled(os.path.join( + outdir, f"raw_m{DEV_M}_offline_metrics.json", + )) + return results + + +def run(args): + device = args.device + outdir = os.path.abspath(args.outdir) + if os.path.exists(outdir): + raise FileExistsError(outdir) + os.makedirs(outdir) + anchor_payload = torch.load( + args.anchor, map_location="cpu", weights_only=False, + ) + anchor_shard = AnchorShard(anchor_payload) + r1_baseline = m10_pooled(args.r1_dev_metrics) + combos = ( + [dict(E=E, lr=lr, anchor=am) + for E, lr, am in itertools.product(E_GRID, LR_GRID, ANCHOR_GRID)] + if args.mode == "develop" + else [dict(E=int(args.fixed_E), lr=float(args.fixed_lr), + anchor=float(args.fixed_anchor))] + ) + opts = ContOpts() + accepted_path = os.path.abspath(args.r1_checkpoint) + accepted_sha = OS.sha256_file(accepted_path) + previous_shard = OS.ExecutedRoundShard.load(args.previous_shard) + template_policy, _ = GPS.load_sfm_policy(accepted_path, device=device) + history = [] + with ProcessPoolExecutor(max_workers=args.verifier_workers) as executor: + for round_k in range(1, int(args.rounds) + 1): + start = time.perf_counter() + round_index = 1 + round_k + accepted_policy, _ = GPS.load_sfm_policy( + accepted_path, device=device, + ) + shard, recovery, gather_info = gather_round( + accepted_policy, previous_shard, round_index, opts, device, + executor, + ) + shard.save(os.path.join( + outdir, "round_shards", f"cont_{round_k:02d}.pt", + )) + candidates = [] + for combo in combos: + name = ( + f"E{combo['E']:02d}_lr{combo['lr']:.0e}" + f"_a{str(combo['anchor']).replace('.', 'p')}" + ).replace("-", "m") + policy = copy.deepcopy(template_policy) + policy.load_state_dict(accepted_policy.state_dict()) + BS.configure_expansion_trainability(policy) + optimizer = torch.optim.Adam( + [p for p in policy.parameters() if p.requires_grad], + lr=combo["lr"], + ) + populations = build_positive_populations( + shard, recovery, anchor_shard, combo["anchor"], + ) + negatives = [(shard, row) for row in shard.Dminus] + replay_log = continuation_replay( + policy, optimizer, populations, negatives, + epochs=combo["E"], batch=BATCH, + seed=SEED + round_k * 1_000_003, device=device, + ) + probe = probe_diagnostics( + policy, accepted_policy, anchor_shard, device, + ) + ckpt = os.path.join( + outdir, f"round_{round_k:02d}", f"{name}.pt", + ) + os.makedirs(os.path.dirname(ckpt), exist_ok=True) + BX._save_checkpoint(policy, ckpt, dict( + round=round_k, combo=combo, accepted_parent=accepted_path, + accepted_parent_sha256=accepted_sha, + )) + candidates.append(dict( + name=name, combo=combo, checkpoint=ckpt, + replay=replay_log, probe=probe, + )) + del policy, optimizer + specs = [ + (c["checkpoint"], os.path.join( + outdir, f"round_{round_k:02d}", f"eval_{c['name']}", + )) + for c in candidates + ] + metrics = evaluate_checkpoints( + specs, cache_dir=os.path.join(outdir, "dev_cache"), + workers_each=args.eval_workers, wave=args.eval_wave, + gpu=args.gpu_index, + ) + for candidate in candidates: + candidate["m10"] = metrics[candidate["checkpoint"]] + ok, checks = admissible(candidate["m10"], r1_baseline) + candidate["admissible"] = ok + candidate["admissibility_checks"] = checks + admissible_rows = [c for c in candidates if c["admissible"]] + selected = ( + min(admissible_rows, key=lambda c: lex_key(c["m10"])) + if admissible_rows else None + ) + record = dict( + continuation_round=round_k, + gather_round_index=round_index, + gather=gather_info, + shard=dict(D=len(shard.D), Dplus=len(shard.Dplus), + Dminus=len(shard.Dminus)), + r1_baseline=r1_baseline, + candidates=[{ + k: c[k] for k in ( + "name", "combo", "checkpoint", "replay", "probe", + "m10", "admissible", "admissibility_checks", + ) + } for c in candidates], + n_admissible=len(admissible_rows), + selected=None if selected is None else dict( + name=selected["name"], combo=selected["combo"], + checkpoint=selected["checkpoint"], + m10=selected["m10"], + ), + wall_seconds=time.perf_counter() - start, + ) + history.append(record) + with open(os.path.join(outdir, "metrics.jsonl"), "a") as stream: + stream.write(json.dumps(record, allow_nan=False) + "\n") + print(json.dumps(dict( + round=round_k, + admissible=len(admissible_rows), + selected=None if selected is None else selected["name"], + selected_m10=None if selected is None else { + k: selected["m10"][k] + for k in ("SR", "CR", "Validity", "clearance", "time") + }, + wall=record["wall_seconds"], + )), flush=True) + if selected is None: + print("STOP: no admissible candidate", flush=True) + break + accepted_path = selected["checkpoint"] + accepted_sha = OS.sha256_file(accepted_path) + previous_shard = shard + OE._write_json(os.path.join(outdir, "COMPLETE.json"), dict( + status="R1_CONTINUATION_COMPLETE", + mode=args.mode, + rounds_run=len(history), + rounds_accepted=sum(1 for r in history if r["selected"]), + final_accepted_checkpoint=accepted_path, + final_accepted_sha256=accepted_sha, + r1_dev_baseline=r1_baseline, + combos=combos, + history=history, + )) + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--mode", choices=("develop", "confirm"), + required=True) + parser.add_argument("--outdir", required=True) + parser.add_argument("--r1-checkpoint", required=True) + parser.add_argument("--previous-shard", required=True, + help="archived B9 round-1 shard (GP chain seed)") + parser.add_argument("--anchor", required=True) + parser.add_argument("--r1-dev-metrics", required=True) + parser.add_argument("--rounds", type=int, default=10) + parser.add_argument("--fixed-E", type=int) + parser.add_argument("--fixed-lr", type=float) + parser.add_argument("--fixed-anchor", type=float) + parser.add_argument("--verifier-workers", type=int, default=14) + parser.add_argument("--eval-workers", type=int, default=12) + parser.add_argument("--eval-wave", type=int, default=8) + parser.add_argument("--gpu-index", type=int, default=3) + parser.add_argument("--device", default="cuda:0") + args = parser.parse_args(argv) + if args.mode == "confirm" and None in ( + args.fixed_E, args.fixed_lr, args.fixed_anchor, + ): + raise SystemExit("confirm mode requires --fixed-E/--fixed-lr/--fixed-anchor") + run(args) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/claude_corrected_distill.py b/overnight_run_07_12_sfm/claude_corrected_distill.py new file mode 100644 index 0000000..9567c54 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_corrected_distill.py @@ -0,0 +1,230 @@ +"""Corrected ordinary-B1 replay and corrected privileged-teacher objective. + +Fixes the two implementation errors named in the study directive: + +A. Ordinary B1 uses the ORIGINAL online semantics via ``BS.signed_update``: + one full-dataset accumulated gradient and exactly one Adam step per replay + epoch, over a fail-closed W=2 window holding BOTH declared recent shards. + +B. The teacher optimizes exactly L_T = sum_i m_i L_CFM(i) for normalized + hierarchy masses (sum m_i = 1): each chunk of size b passes + ``weights = b * m_i`` to the production ``FlowPolicy.cfm_loss`` (which + multiplies then means) and the returned chunk losses are SUMMED — never + multiplied by a chunk mass again. One accumulated Adam step per epoch; + deterministic per-chunk CFM noise via ``torch.manual_seed``. + +The module also provides the label-based metrics reader (the previous +``records[0]`` reader silently returned r0 as the r1 baseline) and the +fixed-context output-RMSE probe. +""" +from __future__ import annotations + +import numpy as np +import torch + +import _paths # noqa: F401 +import sfm_b1_eval as BE +import sfm_b1_store as BS + + +class CorrectedRecent: + """Fail-closed W=2 window over exactly two distinct executed shards.""" + + def __init__(self, shards): + shards = list(shards) + if len(shards) != 2: + raise RuntimeError( + f"W=2 requires BOTH declared recent shards; got {len(shards)} " + "— aborting instead of silently using W=1" + ) + rounds = [int(shard.round_i) for shard in shards] + if len(set(rounds)) != 2: + raise RuntimeError(f"W=2 shards must be distinct rounds: {rounds}") + for shard in shards: + if not shard.Dplus and not shard.Dminus: + raise RuntimeError( + f"declared recent shard round {shard.round_i} is empty" + ) + self._rounds = sorted(shards, key=lambda s: int(s.round_i)) + self.window = 2 + + @property + def rounds(self): + return list(self._rounds) + + def positive_records(self): + return [(s, row) for s in self._rounds for row in s.Dplus] + + def negative_records(self): + return [(s, row) for s in self._rounds for row in s.Dminus] + + +def corrected_ordinary_epochs(policy, optimizer, recent, *, alpha, epochs, + batch, device, seed): + """One ``BS.signed_update`` call (= one Adam step) per epoch.""" + results = [] + for epoch in range(int(epochs)): + result = BS.signed_update( + policy, optimizer, recent, alpha=float(alpha), batch=int(batch), + device=device, seed=int(seed) + epoch, + ) + if int(result["optimizer_steps"]) != 1: + raise RuntimeError( + "corrected ordinary epoch must take exactly one Adam step, " + f"observed {result['optimizer_steps']}" + ) + results.append({ + key: result[key] for key in ( + "path", "rho", "positive_norm", "negative_norm", + "positive_loss", "negative_loss", "positive_eligible", + "negative_eligible", "optimizer_steps", + ) + }) + return dict(epochs=int(epochs), adam_steps=len(results), per_epoch=results) + + +def normalized_teacher_mass(records): + """Hierarchy masses over (holder, row) records, verified to sum to 1.""" + mass, accounting = BS.hierarchy_mass(records) + total = sum(mass.values()) + if not np.isclose(total, 1.0, atol=1e-9): + raise RuntimeError(f"teacher hierarchy mass sums to {total}, not 1") + return mass, accounting + + +def _chunks(sequence, size): + for start in range(0, len(sequence), int(size)): + yield start, sequence[start:start + int(size)] + + +def teacher_epoch_loss(policy, records, mass, *, batch, device, seed, + backward): + """SUM of chunk losses with weights = b * m_i (exactly sum_i m_i L_i).""" + total = 0.0 + for start, values in _chunks(records, batch): + grid, low, hist, controls = BS._tensor_batch(values, device) + context = policy.ctx_from(grid, low, hist) + weights = torch.as_tensor([ + len(values) * mass[(id(holder), int(row["query_id"]))] + for holder, row in values + ], dtype=controls.dtype, device=device) + torch.manual_seed(int(seed) + start) + chunk_loss = policy.cfm_loss(controls, context, weights=weights) + if not bool(torch.isfinite(chunk_loss)): + raise FloatingPointError("non-finite teacher chunk loss") + if backward: + chunk_loss.backward() + total += float(chunk_loss.detach()) + return total + + +def corrected_teacher_epochs(policy, optimizer, records, *, epochs, batch, + device, seed): + """One accumulated Adam step per epoch on the exact L_T objective.""" + mass, accounting = normalized_teacher_mass(records) + policy.train() + losses = [] + for epoch in range(int(epochs)): + optimizer.zero_grad(set_to_none=True) + loss = teacher_epoch_loss( + policy, records, mass, batch=batch, device=device, + seed=int(seed) + epoch * 1_000_003, backward=True, + ) + optimizer.step() + losses.append(loss) + policy.eval() + return dict( + epochs=int(epochs), adam_steps=int(epochs), losses=losses, + records=len(records), mass_accounting_gamma=accounting["gamma"], + ) + + +def replayed_chunk_noise(controls_shape, device, seed): + """Replay cfm_loss's internal draws (x0 then tau) for one chunk (CPU).""" + b, d = controls_shape + torch.manual_seed(int(seed)) + x0 = torch.randn(b, d, device=device) + tau = torch.rand(b, device=device).clamp(1e-4, 1.0) + return x0, tau + + +def direct_teacher_loss(policy, records, mass, *, batch, device, seed): + """Independent sum_i m_i L_i using the replayed per-chunk noise.""" + total = None + for start, values in _chunks(records, batch): + grid, low, hist, controls = BS._tensor_batch(values, device) + context = policy.ctx_from(grid, low, hist) + b = len(values) + x1 = (controls / policy.u_max).reshape(b, policy.d) + x0, tau = replayed_chunk_noise((b, policy.d), device, + int(seed) + start) + x_tau = (1 - tau)[:, None] * x0 + tau[:, None] * x1 + prediction = policy.forward( + x_tau, tau, policy._expand_ctx(context, b), + ) + per = ((prediction - (x1 - x0)) ** 2).mean(dim=1) + masses = torch.as_tensor([ + mass[(id(holder), int(row["query_id"]))] for holder, row in values + ], dtype=per.dtype, device=device) + contribution = (per * masses).sum() + total = contribution if total is None else total + contribution + return total + + +@torch.no_grad() +def fixed_context_rmse(policy_a, policy_b, probes, *, device, seed): + """First-action and full-H10 RMSE between two policies on fixed probes.""" + first, full = [], [] + for index, probe in enumerate(probes): + hp10 = torch.as_tensor(probe["hp10"], device=device)[None].float() + low = torch.as_tensor(probe["low5"], device=device)[None].float() + hist = torch.as_tensor(probe["hist"], device=device)[None].float() + generator = np.random.default_rng(int(seed) + index) + latents = torch.as_tensor(generator.standard_normal( + (8, int(policy_a.d)), dtype=np.float32, + ), device=device) + windows = [] + for policy in (policy_a, policy_b): + ctx = policy.ctx_from(hp10, low, hist) + windows.append(BE.integrate_latents( + policy, latents, ctx.repeat_interleave(8, dim=0), nfe=8, + ).reshape(8, 10, 2)) + delta = (windows[0] - windows[1]) + first.append(float(delta[:, 0].square().mean().sqrt())) + full.append(float(delta.square().mean().sqrt())) + return dict( + probes=len(probes), + first_action_rmse=float(np.mean(first)), + h10_window_rmse=float(np.mean(full)), + ) + + +def pooled_by_label(path, label): + """Label-based reader (fixes the records[0]-as-r1 bug).""" + import json + + with open(path) as stream: + payload = json.load(stream) + matches = [r for r in payload["records"] if r["label"] == str(label)] + if len(matches) != 1: + raise KeyError( + f"{path}: expected exactly one record labeled {label!r}, " + f"found {len(matches)}" + ) + cell = matches[0]["cell"]["summary"] + p = cell["pooled"] + + def _gamma(g): + c = cell["per_gamma"][g] + return dict( + clearance=c["successful_clearance"]["mean"], + time=c["successful_time_to_goal"]["mean"], + ) + + return dict( + SR=float(p["SR"]), CR=float(p["CR"]), timeout=float(p["timeout"]), + Validity=float(p["Validity"]["mean"]), + clearance=p["successful_clearance"]["mean"], + time=p["successful_time_to_goal"]["mean"], + g01=_gamma("0.1"), g10=_gamma("1.0"), + ) diff --git a/overnight_run_07_12_sfm/claude_corrected_study.py b/overnight_run_07_12_sfm/claude_corrected_study.py new file mode 100644 index 0000000..5e4c3af --- /dev/null +++ b/overnight_run_07_12_sfm/claude_corrected_study.py @@ -0,0 +1,346 @@ +"""Phase 4: bounded causal study with corrected replay and corrected teacher. + +Arms (identical seeds, latents, ordering, evaluation bank): + A. immutable r1, no update; + B. corrected ordinary B1 only (alpha=.01, ONE accumulated epoch = ONE Adam + step) on the fail-closed W=2 window {archived B9 round-1 shard, a new + round-2 shard gathered from r1 with the unchanged protocol}; + C. corrected teacher only (applied to r1 BEFORE any ordinary continuation): + lr in {1e-5, 3e-5} x epochs in {1, 4}; + D. best teacher dose (M10 lexicographic among liveness-preserving rows) + followed by B's exact single corrected ordinary epoch. + +Every candidate is evaluated on the fixed raw temperature-1 M10/gamma bank. +No hard-gate termination: the complete Pareto table is reported; promotion to +the disjoint M50 follows the pre-registered non-domination + catastrophic- +liveness rule. All required diagnostics are logged per arm. +""" +from __future__ import annotations + +import argparse +from concurrent.futures import ProcessPoolExecutor +import itertools +import json +import os + +import numpy as np +import torch + +import _paths # noqa: F401 +import claude_continuation as CC +import claude_corrected_distill as CD +import claude_teacher_rounds as TRND +import grid_policy_sfm as GPS +import sfm_b1_expand as BX +import sfm_b1_offline_store as OS +import sfm_b1_r2_alpha_replay as R2 +import sfm_b1_store as BS + +TEACHER_DOSES = tuple( + dict(lr=lr, epochs=epochs) + for lr, epochs in itertools.product((1e-5, 3e-5), (1, 4)) +) +ORDINARY = dict(alpha=0.01, epochs=1, batch=128) +CATASTROPHIC_SR = 0.15 +CATASTROPHIC_TIMEOUT = 0.15 + + +def _teacher_records(buffer_payload): + """Flatten the corrected D_MPC into (holder, row) records with contexts.""" + + class _Holder: + def __init__(self): + self.round_i = 2 + self.contexts = [] + self.windows = [] + + holder = _Holder() + for lineage in buffer_payload["lineages"]: + for window in lineage["windows"]: + context = window["context"] + context_id = len(holder.contexts) + holder.contexts.append(dict( + context_id=context_id, round=2, + scenario_id=int(lineage["episode"]), + gamma=float(lineage["gamma"]), step=int(window["start"]), + state=context["state"], hp10=context["hp10"], + low5=context["low5"], hist=context["hist"], + ped_xy=context["ped_xy"], ped_vel=context["ped_vel"], + )) + holder.windows.append(dict( + window_id=context_id, query_id=context_id, + context_id=context_id, + controls=np.asarray(window["controls"], np.float32), y=1, + )) + return holder, [(holder, row) for row in holder.windows] + + +def _fresh_policy(checkpoint, device): + policy, _ = GPS.load_sfm_policy(checkpoint, device=device) + BS.configure_expansion_trainability(policy) + return policy + + +def _probes(teacher_holder, per_gamma=5): + by_gamma = {} + for context in teacher_holder.contexts: + by_gamma.setdefault(round(context["gamma"], 8), []).append(context) + probes = [] + for gamma in sorted(by_gamma): + probes.extend(by_gamma[gamma][:per_gamma]) + return probes + + +def _gradient_cosines(policy, recent, teacher_records, mass, *, device, + seed): + policy.train() + + def _grad(records, weights_fn): + policy.zero_grad(set_to_none=True) + for start in range(0, len(records), 128): + values = records[start:start + 128] + grid, low, hist, controls = BS._tensor_batch(values, device) + context = policy.ctx_from(grid, low, hist) + torch.manual_seed(int(seed) + start) + loss = policy.cfm_loss( + controls, context, weights=weights_fn(values), + ) + loss.backward() + return { + name: parameter.grad.detach().clone() + for name, parameter in policy.named_parameters() + if parameter.requires_grad and parameter.grad is not None + } + + def _uniform(values): + return None + + positives = recent.positive_records()[:256] + negatives = recent.negative_records()[:256] + teachers = teacher_records[:256] + teacher_gradient = _grad( + teachers, + lambda values: torch.as_tensor( + [len(values) * mass[(id(h), int(r["query_id"]))] + for h, r in values], dtype=torch.float32, device=device), + ) + positive_gradient = _grad(positives, _uniform) + negative_gradient = _grad(negatives, _uniform) + policy.zero_grad(set_to_none=True) + policy.eval() + + def _cos(a, b): + common = sorted(set(a) & set(b)) + num = sum(float((a[n].double() * b[n].double()).sum()) for n in common) + na = sum(float(a[n].double().square().sum()) for n in common) ** 0.5 + nb = sum(float(b[n].double().square().sum()) for n in common) ** 0.5 + return num / (na * nb) if na and nb else None + + return dict( + cos_teacher_positive=_cos(teacher_gradient, positive_gradient), + cos_teacher_negative=_cos(teacher_gradient, negative_gradient), + cos_positive_negative=_cos(positive_gradient, negative_gradient), + ) + + +def run(args): + device = args.device + outdir = os.path.abspath(args.outdir) + if os.path.exists(outdir): + raise FileExistsError(outdir) + os.makedirs(outdir) + r1_path = os.path.abspath(args.r1_checkpoint) + buffer_payload = torch.load( + args.teacher_buffer, map_location="cpu", weights_only=False, + ) + teacher_holder, teacher_records = _teacher_records(buffer_payload) + mass, mass_accounting = CD.normalized_teacher_mass(teacher_records) + probes = _probes(teacher_holder) + shard_b9 = OS.ExecutedRoundShard.load(args.b9_shard) + r1_policy = _fresh_policy(r1_path, device) + + # one new round-2 shard gathered from r1 with the unchanged protocol + opts = CC.ContOpts() + with ProcessPoolExecutor(max_workers=args.verifier_workers) as executor: + new_shard, gather_info = TRND._gather( + r1_policy, shard_b9, 2, opts, device, executor, + ) + new_shard.save(os.path.join(outdir, "round_shards", "round_02.pt")) + recent = CD.CorrectedRecent([shard_b9, new_shard]) + cosines_at_r1 = _gradient_cosines( + _fresh_policy(r1_path, device), recent, teacher_records, mass, + device=device, seed=20260728, + ) + + arms = [] + + def _register(name, policy, training_log): + checkpoint = os.path.join(outdir, "arms", f"{name}.pt") + os.makedirs(os.path.dirname(checkpoint), exist_ok=True) + BX._save_checkpoint(policy, checkpoint, dict(arm=name)) + drift = R2._module_relative_drift( + R2._module_snapshot(r1_policy), R2._module_snapshot(policy), + ) + rmse = CD.fixed_context_rmse( + r1_policy, policy, probes, device=device, seed=20260729, + ) + arms.append(dict( + name=name, checkpoint=checkpoint, training=training_log, + parameter_drift=drift, output_rmse=rmse, + )) + del policy + torch.cuda.empty_cache() + + arms.append(dict( + name="A_r1", checkpoint=r1_path, + training=dict(adam_steps=0), + parameter_drift={}, output_rmse=dict( + first_action_rmse=0.0, h10_window_rmse=0.0, + ), + )) + + policy = _fresh_policy(r1_path, device) + optimizer = torch.optim.Adam( + [p for p in policy.parameters() if p.requires_grad], lr=1e-4, + ) + ordinary_log = CD.corrected_ordinary_epochs( + policy, optimizer, recent, alpha=ORDINARY["alpha"], + epochs=ORDINARY["epochs"], batch=ORDINARY["batch"], device=device, + seed=20260730, + ) + _register("B_ordinary_e1", policy, ordinary_log) + + for dose in TEACHER_DOSES: + name = f"C_teacher_lr{dose['lr']:g}_ep{dose['epochs']}".replace( + "-", "m", + ) + policy = _fresh_policy(r1_path, device) + optimizer = torch.optim.Adam( + [p for p in policy.parameters() if p.requires_grad], + lr=dose["lr"], + ) + teacher_log = CD.corrected_teacher_epochs( + policy, optimizer, teacher_records, epochs=dose["epochs"], + batch=128, device=device, seed=20260731, + ) + _register(name, policy, teacher_log) + + specs = [ + (arm["checkpoint"], os.path.join(outdir, f"eval_{arm['name']}")) + for arm in arms + ] + metrics = CC.evaluate_checkpoints( + specs, cache_dir=os.path.join(outdir, "m10_cache"), + workers_each=args.eval_workers, wave=args.eval_wave, + gpu=args.gpu_index, + ) + for arm in arms: + arm["m10"] = metrics[arm["checkpoint"]] + r1_m10 = next(a for a in arms if a["name"] == "A_r1")["m10"] + + def _liveness_ok(m): + return (m["SR"] >= r1_m10["SR"] - CATASTROPHIC_SR + and m["timeout"] <= r1_m10["timeout"] + CATASTROPHIC_TIMEOUT) + + teacher_rows = [a for a in arms if a["name"].startswith("C_")] + live_teachers = [a for a in teacher_rows if _liveness_ok(a["m10"])] + best_pool = live_teachers if live_teachers else teacher_rows + best_teacher = min(best_pool, key=lambda a: CC.lex_key(a["m10"])) + + policy = _fresh_policy(best_teacher["checkpoint"], device) + optimizer = torch.optim.Adam( + [p for p in policy.parameters() if p.requires_grad], lr=1e-4, + ) + d_log = CD.corrected_ordinary_epochs( + policy, optimizer, recent, alpha=ORDINARY["alpha"], + epochs=ORDINARY["epochs"], batch=ORDINARY["batch"], device=device, + seed=20260730, + ) + name = f"D_{best_teacher['name']}_then_ordinary" + _register(name, policy, dict( + teacher=best_teacher["training"], ordinary=d_log, + adam_steps=best_teacher["training"]["adam_steps"] + + d_log["adam_steps"], + )) + d_arm = arms[-1] + d_metrics = CC.evaluate_checkpoints( + [(d_arm["checkpoint"], + os.path.join(outdir, f"eval_{d_arm['name']}"))], + cache_dir=os.path.join(outdir, "m10_cache"), + workers_each=args.eval_workers, wave=1, gpu=args.gpu_index, + ) + d_arm["m10"] = d_metrics[d_arm["checkpoint"]] + + def _dominated(row, others): + m = row["m10"] + for other in others: + if other is row: + continue + o = other["m10"] + better_eq = (o["CR"] <= m["CR"] + and o["Validity"] >= m["Validity"] + and (o["clearance"] or 0) >= (m["clearance"] or 0)) + strictly = (o["CR"] < m["CR"] + or o["Validity"] > m["Validity"] + or (o["clearance"] or 0) > (m["clearance"] or 0)) + if better_eq and strictly: + return True + return False + + for arm in arms: + m = arm["m10"] + arm["liveness_ok"] = _liveness_ok(m) + arm["improves_safety"] = ( + m["CR"] < r1_m10["CR"] or m["Validity"] > r1_m10["Validity"] + or (m["clearance"] or 0) > (r1_m10["clearance"] or 0) + ) + for arm in arms: + arm["non_dominated"] = not _dominated(arm, arms) + arm["promoted"] = bool( + arm["name"] != "A_r1" and arm["non_dominated"] + and arm["improves_safety"] and arm["liveness_ok"] + ) + + report = dict( + status="CORRECTED_CAUSAL_STUDY_M10_COMPLETE", + teacher_buffer=os.path.abspath(args.teacher_buffer), + teacher_buffer_sha256=OS.sha256_file(args.teacher_buffer), + teacher_records=len(teacher_records), + teacher_mass_gamma=mass_accounting["gamma"], + gather=gather_info, + w2_rounds=[int(s.round_i) for s in recent.rounds], + gradient_cosines_at_r1=cosines_at_r1, + r1_m10=r1_m10, + best_teacher=best_teacher["name"], + arms=[{k: v for k, v in arm.items()} for arm in arms], + promotion_rule=( + "non-dominated on (CR,-Validity,-clearance), improves at least " + "one safety metric vs r1, SR >= r1-0.15, timeout <= r1+0.15" + ), + ) + with open(os.path.join(outdir, "M10_STUDY.json"), "w") as stream: + json.dump(report, stream, indent=1, allow_nan=False, default=float) + print(json.dumps(dict( + best_teacher=best_teacher["name"], + promoted=[a["name"] for a in arms if a.get("promoted")], + r1={k: r1_m10[k] for k in ("SR", "CR", "Validity", "clearance")}, + ), allow_nan=False, default=float), flush=True) + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--outdir", required=True) + parser.add_argument("--r1-checkpoint", required=True) + parser.add_argument("--b9-shard", required=True) + parser.add_argument("--teacher-buffer", required=True) + parser.add_argument("--verifier-workers", type=int, default=14) + parser.add_argument("--eval-workers", type=int, default=12) + parser.add_argument("--eval-wave", type=int, default=6) + parser.add_argument("--gpu-index", type=int, default=3) + parser.add_argument("--device", default="cuda:0") + args = parser.parse_args(argv) + run(args) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/claude_kazuki_eval.py b/overnight_run_07_12_sfm/claude_kazuki_eval.py new file mode 100644 index 0000000..46e6961 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_kazuki_eval.py @@ -0,0 +1,171 @@ +"""Locked Kazuki comparator on a fixed M/gamma bank with executed-window Validity. + +Additive evaluation-only script: it changes nothing in the B1 pipeline. The +comparator is the existing generate-guide-refine ``kazuki_sfm_deploy`` with the +locked configuration (safe_coefs=(0.3,), goal_coef=0.5, zero gamma spans, +sample_seed=700000) on the same Hp10 pretrained prior. Episodes are seeded with +``SS.make_humans(episode, 0, n_ped, speed_range)`` exactly as the raw offline +evaluator, so a shared ep0 gives the same pedestrian bank. Validity is the +identical terminal-truncated executed-window metric from +``sfm_b1_offline_eval._verify_executed_episode``. +""" +from __future__ import annotations + +import argparse +from concurrent.futures import ProcessPoolExecutor +import json +import multiprocessing as mp +import os + +import numpy as np + +import _paths # noqa: F401 + + +LOCKED = dict(safe_coef=0.3, goal_coef=0.5, sample_seed=700_000) + + +def _rollout_gamma(payload): + """Worker: run every episode of one gamma cell and attach validity.""" + (checkpoint, scene_profile, ep0, m_per_gamma, gamma, device) = payload + import torch # noqa: F401 (worker-local import keeps spawn cheap to reason about) + import grid_policy_sfm as GPS + import sfm_b1_offline_eval as OE + import sfm_kazuki as KZ + import sfm_protocol as SP + import sfm_scene as SS + + environment = SS.scene_profile(scene_profile) + policy, _ = GPS.load_sfm_policy(checkpoint, device=device) + policy.eval() + config = KZ.KazukiConfig( + safe_coefs=(float(LOCKED["safe_coef"]),), + goal_coef=float(LOCKED["goal_coef"]), + ).validate() + rows = [] + for episode in range(int(ep0), int(ep0) + int(m_per_gamma)): + rollout = KZ.kazuki_sfm_deploy( + policy, episode, float(gamma), cfg=config, + n_ped=environment["n_ped"], T=SP.T, device=device, + ped_speed_range=tuple(environment["ped_speed_range"]), + sample_seed=int(LOCKED["sample_seed"]), collect_diagnostics=False, + ) + success = bool(rollout["success"]) + collision = bool(rollout["collision"]) + steps = int(rollout["steps"]) + row = { + "episode": int(episode), + "gamma": float(gamma), + "status": ( + "success" if success + else "collision" if collision else "timeout" + ), + "success": success, + "collision": collision, + "timeout": bool(not success and not collision), + "steps": steps, + "time_to_goal": steps * SS.DT if success else None, + "min_clearance": float(rollout["min_clear"]), + "successful_clearance": ( + float(rollout["min_clear"]) if success else None + ), + "states": np.asarray(rollout["states"], np.float32), + "controls": np.asarray(rollout["controls"], np.float32), + "ped_xy": np.asarray(rollout["peds"], np.float32), + "ped_vel": np.asarray(rollout["ped_vels"], np.float32), + } + validity = OE._verify_executed_episode(row) + for key in ("states", "controls", "ped_xy", "ped_vel"): + row.pop(key) + row.update(validity) + rows.append(row) + return rows + + +def run(args) -> dict: + import sfm_b1_offline_eval as OE + import sfm_kazuki as KZ + import sfm_protocol as SP + import sfm_scene as SS + + output_dir = os.path.abspath(args.output_dir) + os.makedirs(output_dir, exist_ok=True) + checkpoint = os.path.abspath(args.checkpoint) + checkpoint_sha = OE._sha256_file(checkpoint) + payloads = [ + ( + checkpoint, args.scene_profile, int(args.ep0), + int(args.m_per_gamma), float(gamma), args.device, + ) + for gamma in SP.GAMMAS + ] + context = mp.get_context("spawn") + with ProcessPoolExecutor( + max_workers=min(len(payloads), int(args.workers)), + mp_context=context, + ) as executor: + cell_rows = list(executor.map(_rollout_gamma, payloads)) + rows = [row for cell in cell_rows for row in cell] + summary = OE.summarize( + rows, seed=int(args.ep0) + int(checkpoint_sha[:8], 16) % 100_000, + ) + OE._assert_zero_verifier_errors(summary) + result = { + "status": "CLAUDE_KAZUKI_FIXED_BANK_COMPLETE", + "method": "default Kazuki generate-guide-refine (locked)", + "kazuki_config": dict(LOCKED), + "kazuki_config_full": { + key: (list(value) if isinstance(value, tuple) else value) + for key, value in vars(KZ.KazukiConfig( + safe_coefs=(float(LOCKED["safe_coef"]),), + goal_coef=float(LOCKED["goal_coef"]), + ).validate()).items() + if isinstance(value, (int, float, str, bool, tuple, type(None))) + }, + "checkpoint": checkpoint, + "checkpoint_sha256": checkpoint_sha, + "scene_profile": args.scene_profile, + "environment": SS.scene_profile(args.scene_profile), + "bank": { + "ep0": int(args.ep0), + "M_per_gamma": int(args.m_per_gamma), + "same_scenario_ids_for_every_gamma": True, + "pedestrian_seeding": "SS.make_humans(episode, 0, n_ped, speed_range)", + }, + "summary": summary, + "rows": rows, + "metric_semantics": { + "Validity": ( + "identical executed sliding-window metric as " + "sfm_b1_offline_eval (H_t=min(10,N_tau-t), exact GREEN verifier)" + ), + "comparator_semantics": ( + "learned prior plus reward guidance and MPPI refinement; " + "no retuning; no external shield or fallback" + ), + }, + } + path = os.path.join( + output_dir, f"kazuki_m{int(args.m_per_gamma)}_metrics.json" + ) + OE._write_json(path, result) + print(path) + return result + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--checkpoint", required=True) + parser.add_argument( + "--scene-profile", default="double_density_velocity_ood", + ) + parser.add_argument("--ep0", type=int, required=True) + parser.add_argument("--m-per-gamma", type=int, required=True) + parser.add_argument("--device", default="cuda:0") + parser.add_argument("--workers", type=int, default=7) + parser.add_argument("--output-dir", required=True) + return parser + + +if __name__ == "__main__": + run(build_parser().parse_args()) diff --git a/overnight_run_07_12_sfm/claude_mechanism_viz.py b/overnight_run_07_12_sfm/claude_mechanism_viz.py new file mode 100644 index 0000000..2539d22 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_mechanism_viz.py @@ -0,0 +1,498 @@ +"""Mechanism visualizations for the SFM recipe study (diagnostic-only). + +``snapshot``: for one stored hard context (from an archived ExecutedRoundShard) +render, per checkpoint, the K=16 flow candidates regenerated with the EXACT +keyed gathering latents and labeled by the exact full-H10 verifier; overlay the +deterministic certified recovery escapes (family v1 and, when available, v2); +and show the frozen visual-encoder input (Hp10 polar stack) plus low5/history +conditioning for that context. + +``episode``: closed-loop replay of one (scenario, gamma) gathering lineage +with a given checkpoint (margin-selector B1 semantics, round-1 acquisition +state), recording the trajectory, per-step K-positive counts and NVP flags; +``render-episodes`` overlays two replays (e.g. r0 vs a treated checkpoint). + +These figures are explanatory evidence only; claims use fixed-bank metrics. +""" +from __future__ import annotations + +import argparse +from concurrent.futures import ProcessPoolExecutor +import copy +import json +import os + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import torch + +import _paths # noqa: F401 +import claude_offline_aug as AUG +import grid_policy_sfm as GPS +import sfm_b1_cost as BC +import sfm_b1_eval as BE +import sfm_b1_expand as BX +import sfm_b1_full_episode_audit as FA +import sfm_b1_offline_exec as OE +import sfm_b1_offline_store as OS +import sfm_b1_rbf as BR +import sfm_metrics2 as SM +import sfm_protocol as SP +import sfm_scene as SS + +SEED = 20260724 + + +def _find_context(shard, scenario, gamma, step): + for context in shard.contexts: + if ( + int(context["scenario_id"]) == int(scenario) + and round(float(context["gamma"]), 8) == round(float(gamma), 8) + and int(context["step"]) == int(step) + ): + return context + raise KeyError(f"context not stored: s{scenario} g{gamma} step{step}") + + +def _executed_window(shard, context): + for window in shard.windows: + if int(window["context_id"]) == int(context["context_id"]): + return window + return None + + +@torch.no_grad() +def _k_candidates(policy, context, device): + hp10 = torch.as_tensor(context["hp10"], device=device)[None] + low = torch.as_tensor(context["low5"], device=device)[None] + hist = torch.as_tensor(context["hist"], device=device)[None] + ctx = policy.ctx_from(hp10.float(), low.float(), hist.float()) + generator = np.random.default_rng(OE._keyed_seed( + SEED, 1, int(context["scenario_id"]), + f"{float(context['gamma']):.8f}", int(context["step"]), "K", + )) + x0 = generator.standard_normal((16, int(policy.d)), dtype=np.float32) + windows = BE.integrate_latents( + policy, + torch.as_tensor(x0, device=device), + ctx.repeat_interleave(16, dim=0), + nfe=8, + ).reshape(16, 10, 2).cpu().numpy() + return windows + + +def _verify_many(context, control_sets, workers=8): + tasks = [ + (index, 0, context["state"], controls, context["ped_xy"], + context["ped_vel"], context["gamma"]) + for index, controls in enumerate(control_sets) + ] + with ProcessPoolExecutor(max_workers=workers) as executor: + results = list(executor.map(SM.verify_in_worker, tasks)) + ordered = [None] * len(control_sets) + for index, _, result in results: + ordered[index] = result + return ordered + + +def _scene(axis, context, title): + state = np.asarray(context["state"], np.float32) + ped_xy = np.asarray(context["ped_xy"], np.float32) + ped_vel = np.asarray(context["ped_vel"], np.float32) + prediction = SM.predict_pedestrians(ped_xy, ped_vel, H=10) + for j in range(len(ped_xy)): + axis.add_patch(plt.Circle( + ped_xy[j], SS.R_PED, color="#c2554f", alpha=.75, lw=0, zorder=3, + )) + axis.plot( + prediction[:, j, 0], prediction[:, j, 1], + color="#c2554f", lw=.7, ls=":", alpha=.55, zorder=2, + ) + axis.plot(*SS.GOAL, marker="*", ms=17, color="#e6b422", mec="k", + zorder=6) + axis.plot(state[0], state[1], marker="o", ms=9, color="#1450a3", + mec="k", zorder=6) + axis.annotate( + "", xy=state[:2] + 0.5 * state[2:4], xytext=state[:2], + arrowprops=dict(arrowstyle="->", color="#1450a3", lw=2), zorder=6, + ) + axis.add_patch(plt.Rectangle( + (SS.TASK_LO, SS.TASK_LO), SS.TASK_HI - SS.TASK_LO, + SS.TASK_HI - SS.TASK_LO, fill=False, ec="k", lw=.8, alpha=.6, + )) + axis.add_patch(plt.Circle( + state[:2], SS.R_SENSE, fill=False, ec="#1450a3", lw=.6, ls="--", + alpha=.5, + )) + pad = 2.35 + axis.set_xlim(state[0] - pad, state[0] + pad) + axis.set_ylim(state[1] - pad, state[1] + pad) + axis.set_aspect("equal") + axis.set_title(title, fontsize=11) + axis.grid(alpha=.2) + + +def _draw_windows(axis, context, windows, labels, *, executed=None): + n_pos = 0 + for controls, result in zip(windows, labels): + segment = SM.rollout_positions(context["state"], controls) + positive = bool(result.get("resolved")) and int(result.get("y", 0)) == 1 + n_pos += int(positive) + axis.plot( + segment[:, 0], segment[:, 1], + color="#1f8a4c" if positive else "#b22222", + lw=1.7 if positive else 0.9, + alpha=.95 if positive else .5, + zorder=5 if positive else 4, + ) + if executed is not None: + segment = SM.rollout_positions(context["state"], executed) + axis.plot(segment[:, 0], segment[:, 1], color="#550000", lw=3.2, + alpha=.95, zorder=5.5, label="executed (uncertified)") + return n_pos + + +def snapshot(args): + shard = OS.ExecutedRoundShard.load(args.shard) + context = _find_context(shard, args.scenario, args.gamma, args.step) + executed = _executed_window(shard, context) + specs = [spec.split("=", 1) for spec in args.checkpoints] + families = dict(v1="v1", v2="v2") if hasattr(AUG, "recovery_candidates_v2") \ + else dict(v1="v1") + n_ckpt = len(specs) + n_cols = n_ckpt + len(families) + figure, axes = plt.subplots( + 2, max(n_cols, 3), figsize=(4.9 * max(n_cols, 3), 9.6), + ) + summary = dict( + scenario=int(args.scenario), gamma=float(args.gamma), + step=int(args.step), + ) + + for column, (name, path) in enumerate(specs): + policy, _ = GPS.load_sfm_policy(path, device=args.device) + policy.eval() + windows = _k_candidates(policy, context, args.device) + results = _verify_many(context, list(windows), args.workers) + axis = axes[0][column] + _scene(axis, context, "") + n_pos = _draw_windows( + axis, context, windows, results, + executed=None if executed is None or column else + np.asarray(executed["controls"], np.float32), + ) + axis.set_title( + f"{name}: K=16 flow candidates\n" + f"exact-verifier positives: {n_pos}/16", fontsize=11, + ) + summary[f"K_positive_{name}"] = int(n_pos) + del policy + + for offset, family in enumerate(sorted(families)): + candidates = ( + AUG.recovery_candidates(context["state"]) if family == "v1" + else AUG.recovery_candidates_v2( + context["state"], context["ped_xy"], context["ped_vel"], + ) + ) + scored = [] + for controls, provenance in candidates: + objective = AUG._prefilter(context, controls) + if objective is not None: + scored.append((objective, controls, provenance)) + scored.sort(key=lambda row: (row[0], str(row[2]))) + pool = scored[:AUG.PREVERIFY_CAP] + results = _verify_many(context, [row[1] for row in pool], args.workers) + axis = axes[0][n_ckpt + offset] + _scene(axis, context, "") + certified = 0 + best_drawn = False + for (objective, controls, provenance), result in zip(pool, results): + segment = SM.rollout_positions(context["state"], controls) + ok = bool(result.get("resolved")) and int(result.get("y", 0)) == 1 + if ok: + certified += 1 + axis.plot( + segment[:, 0], segment[:, 1], + color="#0b6fa4" if family == "v1" else "#e07b00", + lw=3.0 if not best_drawn else 1.6, + alpha=.95 if not best_drawn else .7, zorder=5.4, + ) + if not best_drawn: + summary[f"recovery_{family}_best"] = dict( + J=float(objective), generator=provenance, + slack=float(result["diagnostics"]["slack"]), + end_speed=float(np.linalg.norm( + BC.rollout_states( + context["state"], controls[None], + )[0, -1, 2:4].numpy() + )), + ) + best_drawn = True + else: + axis.plot(segment[:, 0], segment[:, 1], color="#888888", + lw=.7, alpha=.4, zorder=3.5) + axis.set_title( + f"certified deterministic recovery ({family})\n" + f"{certified}/{len(pool)} exact-certified", fontsize=11, + ) + summary[f"recovery_{family}_certified"] = int(certified) + summary[f"recovery_{family}_pool"] = int(len(pool)) + + hp10 = np.asarray(context["hp10"], np.float32) + axis = axes[1][0] + image = axis.imshow( + hp10[-1].T, origin="lower", aspect="auto", cmap="RdBu", + vmin=-1, vmax=1, + extent=(-180, 180, 0, SS.R_SENSE), + ) + axis.set_title("frozen encoder input: newest H_P frame\n" + "(clipped nominal polytope, polar)", fontsize=11) + axis.set_xlabel("bearing [deg]") + axis.set_ylabel("range [m]") + plt.colorbar(image, ax=axis, fraction=.04) + for column in range(1, min(3, axes.shape[1])): + axis = axes[1][column] + if column == 1: + mosaic = np.concatenate([hp10[i].T for i in range(10)], axis=1) + axis.imshow( + mosaic, origin="lower", aspect="auto", cmap="RdBu", + vmin=-1, vmax=1, + ) + axis.set_title("Hp10 stack: 10 most recent H_P frames " + "(oldest left)", fontsize=11) + axis.set_xticks([]) + axis.set_yticks([]) + elif column == 2: + axis.axis("off") + low = np.asarray(context["low5"], np.float32) + lines = [ + f"scenario {args.scenario} gamma {args.gamma} " + f"step {args.step}", + f"low5: relgoal=({low[0]:.2f},{low[1]:.2f}) " + f"v=({low[2]:.2f},{low[3]:.2f}) gamma={low[4]:.2f}", + "encoder enc_grid: FROZEN during expansion (SHA-checked)", + "gradient flows into trunk/GRU/enc_low conditioned on", + "these frozen grid features", + ] + if executed is not None: + lines.append( + f"executed window: y={executed['y']} " + f"source={executed['execution_source']}" + ) + for key in sorted(summary): + if key.startswith(("K_positive", "recovery")): + value = summary[key] + if isinstance(value, dict): + value = { + k: (round(v, 3) if isinstance(v, float) else v) + for k, v in value.items() if k != "generator" + } + lines.append(f"{key}: {value}") + axis.text(0.01, 0.98, "\n".join(str(l) for l in lines), + va="top", ha="left", fontsize=9, family="monospace", + transform=axis.transAxes, wrap=True) + for column in range(n_cols, axes.shape[1]): + axes[0][column].axis("off") + for column in range(3, axes.shape[1]): + axes[1][column].axis("off") + figure.suptitle( + f"Hard-context mechanism: s{args.scenario} γ={args.gamma} " + f"step {args.step} (exact verifier everywhere)", fontsize=13, + ) + figure.tight_layout(rect=(0, 0, 1, 0.96)) + figure.savefig(args.out, dpi=170, bbox_inches="tight") + plt.close(figure) + OE._write_json(args.out + ".json", summary) + print(json.dumps({k: v for k, v in summary.items() + if not isinstance(v, dict)}, indent=1)) + + +@torch.no_grad() +def episode(args): + device = args.device + policy, _ = GPS.load_sfm_policy(args.checkpoint, device=device) + policy.eval() + phi_policy = copy.deepcopy(policy).eval() + for parameter in phi_policy.parameters(): + parameter.requires_grad_(False) + cfg = OE.OfflineConfig(alpha=0.0, exposure_epochs=1, rounds=1, smoke=True) + environment = SS.scene_profile(cfg.scene_profile) + replica = BX.Replica( + int(args.scenario), float(args.gamma), + n_ped=environment["n_ped"], + ped_speed_range=tuple(environment["ped_speed_range"]), + ) + gp = BR.RBFGP(float(args.ell), float(cfg.gp_lam)) + beta, _ = OE._calibrate_beta( + phi_policy, gp, [replica], cfg, device, round_i=1, + ) + frames = [] + with ProcessPoolExecutor(max_workers=args.workers) as executor: + for step in range(int(cfg.T)): + live, batch = BX._stack_prepared([replica], device) + if not live: + break + windows, contexts, x0 = OE._keyed_windows( + policy, live, batch, K=cfg.K, round_i=1, step=step, + source="K", seed=cfg.seed, nfe=cfg.nfe, temp=cfg.temp, + ) + raw_windows, _, _ = OE._keyed_windows( + policy, live, batch, K=1, round_i=1, step=step, + source="raw_continuation", seed=cfg.seed, nfe=cfg.nfe, + temp=cfg.temp, + ) + windows_np = windows[0].cpu().numpy() + raw_np = raw_windows[0, 0].cpu().numpy() + features = OE._features_from_x0( + phi_policy, windows, contexts, x0, cfg.phi_s, + ) + generator = torch.Generator(device=features.device) + generator.manual_seed(OE._keyed_seed( + cfg.seed, 1, replica.scenario_id, + f"{replica.gamma:.8f}", step, "acquisition", + )) + selected, trace = gp.sequential_acquire( + features[0], cfg.B, beta, generator=generator, + ) + prepared = replica.prepared + results = _verify_many( + dict(state=prepared["state"], ped_xy=prepared["ped_xy"], + ped_vel=prepared["ped_vel"], gamma=replica.gamma), + list(windows_np), args.workers, + ) + query_rows = [ + dict(candidate_id=int(k), acquisition_step=j, + controls=windows_np[k], result=results[k], mode=None, + sigma=float(trace[j]["chosen_sigma"])) + for j, k in enumerate(map(int, selected)) + if results[k].get("resolved") + ] + chosen = BC.select_admissible( + query_rows, selector="margin", state=prepared["state"], + ped_xy=prepared["ped_xy"], ped_vel=prepared["ped_vel"], + gamma=replica.gamma, + ) + controls = raw_np if chosen is None else np.asarray( + chosen["controls"], np.float32, + ) + frames.append(dict( + step=int(step), + state=prepared["state"].tolist(), + ped_xy=prepared["ped_xy"].tolist(), + K_positive=int(sum( + int(r.get("y", 0)) == 1 for r in results + if r.get("resolved") + )), + NVP=chosen is None, + )) + BX._advance(replica, controls[0]) + FA._post_action_terminal(replica) + OE._finalize_alive([replica]) + payload = dict( + checkpoint=os.path.abspath(args.checkpoint), + scenario=int(args.scenario), gamma=float(args.gamma), + status=replica.status, steps=len(replica.controls), + min_clearance=float(replica.minimum_clearance), + states=[s.tolist() for s in replica.states], + frames=frames, + ) + OE._write_json(args.out, payload) + print(json.dumps(dict(status=replica.status, + steps=len(replica.controls)), indent=1)) + + +def render_episodes(args): + runs = [] + for spec in args.runs: + name, path = spec.split("=", 1) + with open(path) as stream: + runs.append((name, json.load(stream))) + n = len(runs) + figure, axes = plt.subplots(1, n, figsize=(6.4 * n, 6.4)) + if n == 1: + axes = [axes] + for axis, (name, run) in zip(axes, runs): + states = np.asarray(run["states"], np.float32) + nvp_steps = {f["step"] for f in run["frames"] if f["NVP"]} + final = run["frames"][-1] + ped = np.asarray(final["ped_xy"], np.float32) + for j in range(len(ped)): + axis.add_patch(plt.Circle( + ped[j], SS.R_PED, color="#c2554f", alpha=.5, lw=0, + )) + for t in range(len(states) - 1): + color = "#d95f02" if t in nvp_steps else "#1450a3" + axis.plot(states[t:t + 2, 0], states[t:t + 2, 1], color=color, + lw=2.6 if t in nvp_steps else 1.8, zorder=5) + axis.plot(*SS.GOAL, marker="*", ms=17, color="#e6b422", mec="k") + axis.plot(states[0, 0], states[0, 1], marker="s", ms=8, + color="#1450a3", mec="k") + marker = dict(collision="X", success="*", timeout="P")[run["status"]] + axis.plot(states[-1, 0], states[-1, 1], marker=marker, ms=14, + color={"collision": "#b22222", "success": "#1f8a4c", + "timeout": "#888888"}[run["status"]], mec="k", + zorder=7) + axis.add_patch(plt.Rectangle( + (SS.TASK_LO, SS.TASK_LO), SS.TASK_HI - SS.TASK_LO, + SS.TASK_HI - SS.TASK_LO, fill=False, ec="k", lw=.8, alpha=.6, + )) + nvp_count = len(nvp_steps) + axis.set_title( + f"{name}: {run['status']} in {run['steps']} steps\n" + f"NVP steps (orange): {nvp_count}; min clearance " + f"{run['min_clearance']:.3f} m", fontsize=12, + ) + axis.set_aspect("equal") + axis.set_xlim(SS.TASK_LO - .2, SS.TASK_HI + .2) + axis.set_ylim(SS.TASK_LO - .2, SS.TASK_HI + .2) + axis.grid(alpha=.2) + figure.suptitle( + f"Closed-loop gathering lineage s{runs[0][1]['scenario']} " + f"γ={runs[0][1]['gamma']} — final pedestrian frame shown", + fontsize=13, + ) + figure.tight_layout(rect=(0, 0, 1, 0.94)) + figure.savefig(args.out, dpi=170, bbox_inches="tight") + plt.close(figure) + print(args.out) + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + sub = parser.add_subparsers(dest="cmd", required=True) + + s = sub.add_parser("snapshot") + s.add_argument("--shard", required=True) + s.add_argument("--scenario", type=int, required=True) + s.add_argument("--gamma", type=float, required=True) + s.add_argument("--step", type=int, required=True) + s.add_argument("--checkpoints", nargs="+", required=True, + help="NAME=PATH ...") + s.add_argument("--workers", type=int, default=8) + s.add_argument("--device", default="cpu") + s.add_argument("--out", required=True) + + e = sub.add_parser("episode") + e.add_argument("--checkpoint", required=True) + e.add_argument("--scenario", type=int, required=True) + e.add_argument("--gamma", type=float, required=True) + e.add_argument("--ell", type=float, required=True) + e.add_argument("--workers", type=int, default=8) + e.add_argument("--device", default="cuda:0") + e.add_argument("--out", required=True) + + r = sub.add_parser("render-episodes") + r.add_argument("--runs", nargs="+", required=True, help="NAME=JSON ...") + r.add_argument("--out", required=True) + + args = parser.parse_args(argv) + dict(snapshot=snapshot, episode=episode, + render_episodes=render_episodes)[args.cmd.replace("-", "_")](args) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/claude_mpc_exec.py b/overnight_run_07_12_sfm/claude_mpc_exec.py new file mode 100644 index 0000000..b054381 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_mpc_exec.py @@ -0,0 +1,426 @@ +"""Ordinary B1 expansion + dedicated per-round D_MPC+ distillation blocks. + +Per macro-round: +1. ordinary gather (immutable ``sfm_b1_offline_exec.gather_offline_round``) + and ordinary replay (untouched ``sfm_b1_offline_replay.replay``); +2. save the post-ordinary checkpoint (``round_XX_pre_block.pt``); +3. harvest this round's ``D_MPC+`` (``claude_mpc_pool``): privileged Codex + pool ∧ exact full-H10 SOCP, ranked by native SafeMPPI cost, per-round + fresh, gamma/context-balanced through the standard hierarchy mass; +4. audit BEFORE the block, run the dedicated distillation block (its own + Adam, swept lr/epochs), audit AFTER, save ``round_XX.pt``. + +Audits per block: (a) local MPC-context raw sampling — SOCP-positive rate, +mean predicted clearance and goal progress of 16 raw temperature-1 samples +per held audit context, and target recovery (min normalized L2 distance from +samples to the stored MPC target); (b) fixed raw temperature-1 M10/gamma +evaluation (CR, executed-window Validity, successful clearance and time) on +the declared MPC study bank. D_MPC+ never touches D/D+, the GP, or +acquisition; the privileged controller is never used at evaluation. +""" +from __future__ import annotations + +import argparse +from concurrent.futures import ProcessPoolExecutor +import copy +from dataclasses import asdict, dataclass +import json +import os +import time + +import numpy as np +import torch + +import _paths # noqa: F401 +import claude_mpc_pool as MP +import claude_offline_aug as AUG +import grid_policy_sfm as GPS +import sfm_b1_eval as BE +import sfm_b1_expand as BX +import sfm_b1_offline_exec as OE +import sfm_b1_offline_replay as OR +import sfm_b1_offline_store as OS +import sfm_b1_store as BS +import sfm_metrics2 as SM +import sfm_protocol as SP +import sfm_scene as SS + +M10_EP0 = 350_000 +M10_NOISE_SEED = 20_260_733 +AUDIT_CONTEXTS = 64 +AUDIT_SAMPLES = 16 +MAX_HARVEST_CONTEXTS = 400 + + +@dataclass(frozen=True) +class MPCConfig: + lr_dedicated: float + epochs_dedicated: int + selector: str = "margin" + alpha: float = 0.01 + exposure_epochs: int = 10 + lr: float = 1.0e-4 + ess_target: float = 0.5 + rounds: int = 4 + K: int = 16 + B: int = 4 + T: int = 180 + H: int = 10 + batch: int = 128 + nfe: int = 8 + temp: float = 1.0 + phi_s: float = 0.9 + gp_lam: float = OE.GP_LAMBDA + verifier_workers: int = 8 + seed: int = 20260724 + scene_profile: str = OE.SCENE_PROFILE + smoke: bool = False + tag: str = "mpc" + + def validate(self): + if not 0.0 < float(self.lr_dedicated) <= 1.0e-3: + raise ValueError("dedicated lr out of range") + if not 1 <= int(self.epochs_dedicated) <= 32: + raise ValueError("dedicated epochs out of range") + if ( + int(self.K), int(self.B), int(self.T), int(self.H), + int(self.batch), float(self.gp_lam), float(self.temp), + self.scene_profile, int(self.nfe), float(self.phi_s), + ) != (16, 4, 180, 10, 128, OE.GP_LAMBDA, 1.0, OE.SCENE_PROFILE, 8, 0.9): + raise ValueError("immutable offline core changed") + return self + + @property + def arm_name(self): + lr = f"{self.lr_dedicated:.0e}".replace("-", "m") + return f"{self.tag}_lrd{lr}_epd{int(self.epochs_dedicated):02d}" + + +@torch.no_grad() +def local_mpc_audit(policy, shard, records, *, device, executor, seed_tag): + """Raw-sampling audit at gamma-balanced held D_MPC+ contexts.""" + if not records: + return dict(contexts=0) + by_gamma = {} + for record in records: + gamma = round(float(shard.contexts[record["context_id"]]["gamma"]), 8) + by_gamma.setdefault(gamma, []).append(record) + chosen = [] + quota = max(1, AUDIT_CONTEXTS // max(len(by_gamma), 1)) + for gamma in sorted(by_gamma): + rows = sorted(by_gamma[gamma], key=lambda r: ( + r["context_id"], r["rank"], + )) + seen = set() + for row in rows: + if row["context_id"] in seen: + continue + seen.add(row["context_id"]) + chosen.append(row) + if len(seen) >= quota: + break + positive = clearance = progress = recovery = support = 0.0 + for row in chosen: + context = shard.contexts[row["context_id"]] + hp10 = torch.as_tensor(context["hp10"], device=device)[None].float() + low = torch.as_tensor(context["low5"], device=device)[None].float() + hist = torch.as_tensor(context["hist"], device=device)[None].float() + ctx = policy.ctx_from(hp10, low, hist) + generator = np.random.default_rng(OE._keyed_seed( + 20260724, 99, int(context["scenario_id"]), + f"{float(context['gamma']):.8f}", int(context["step"]), + f"mpc_audit_{seed_tag}", + )) + x0 = generator.standard_normal( + (AUDIT_SAMPLES, int(policy.d)), dtype=np.float32, + ) + windows = BE.integrate_latents( + policy, torch.as_tensor(x0, device=device), + ctx.repeat_interleave(AUDIT_SAMPLES, dim=0), nfe=8, + ).reshape(AUDIT_SAMPLES, 10, 2).cpu().numpy() + tasks = [ + (k, 0, context["state"], windows[k], context["ped_xy"], + context["ped_vel"], context["gamma"]) + for k in range(AUDIT_SAMPLES) + ] + results = {k: r for k, _, r in executor.map(SM.verify_in_worker, tasks)} + n_pos = sum( + 1 for r in results.values() + if r.get("resolved") and int(r.get("y", 0)) == 1 + ) + positive += n_pos / AUDIT_SAMPLES + support += float(n_pos > 0) + geometry = [ + AUG._window_geometry(context, windows[k]) + for k in range(AUDIT_SAMPLES) + ] + clearance += float(np.mean([g[0] for g in geometry])) + state = np.asarray(context["state"], np.float32) + goal_now = float(np.linalg.norm(state[:2] - SS.GOAL)) + segs = [SM.rollout_positions(state, windows[k])[-1] + for k in range(AUDIT_SAMPLES)] + progress += float(np.mean([ + goal_now - float(np.linalg.norm(seg - SS.GOAL)) for seg in segs + ])) + target = np.asarray(row["controls"], np.float32) + distances = [ + float(np.linalg.norm(windows[k] - target) / np.sqrt(target.size)) + for k in range(AUDIT_SAMPLES) + ] + recovery += min(distances) + n = max(len(chosen), 1) + return dict( + contexts=len(chosen), + socp_positive_rate=positive / n, + support_fraction=support / n, + mean_sample_clearance=clearance / n, + mean_sample_progress=progress / n, + target_recovery_rmse=recovery / n, + ) + + +def m10_eval(checkpoint, label, *, outdir, cache_dir, workers, device): + import sfm_b1_offline_eval as EV + args = argparse.Namespace( + checkpoints=[checkpoint], labels=[label], + scene_profile=OE.SCENE_PROFILE, ep0=M10_EP0, + noise_seed=M10_NOISE_SEED, m_per_gamma=10, device=device, + workers=int(workers), cache_dir=cache_dir, output_dir=outdir, + ) + result = EV.run(args) + pooled = result["records"][0]["cell"]["summary"]["pooled"] + return dict( + SR=float(pooled["SR"]), CR=float(pooled["CR"]), + timeout=float(pooled["timeout"]), + Validity=float(pooled["Validity"]["mean"]), + clearance=pooled["successful_clearance"]["mean"], + time=pooled["successful_time_to_goal"]["mean"], + per_gamma={ + gamma: dict( + CR=cell["CR"], + Validity=float(cell["Validity"]["mean"]), + clearance=cell["successful_clearance"]["mean"], + time=cell["successful_time_to_goal"]["mean"], + ) + for gamma, cell in + result["records"][0]["cell"]["summary"]["per_gamma"].items() + }, + ) + + +def run(checkpoint, outdir, cfg, *, device): + cfg.validate() + checkpoint = os.path.abspath(checkpoint) + outdir = os.path.abspath(outdir) + checkpoint_sha = OS.sha256_file(checkpoint) + if checkpoint_sha != OE.EXPECTED_CHECKPOINT_SHA256: + raise ValueError("MPC study must start from the exact r0 checkpoint") + if os.path.exists(outdir): + raise FileExistsError(outdir) + os.makedirs(outdir) + environment = SS.scene_profile(cfg.scene_profile) + policy, _ = GPS.load_sfm_policy(checkpoint, device=device) + BS.configure_expansion_trainability(policy) + encoder_sha = BS.module_sha256(policy.enc_grid) + optimizer = torch.optim.Adam( + [p for p in policy.parameters() if p.requires_grad], lr=cfg.lr, + ) + optimizer_dedicated = torch.optim.Adam( + [p for p in policy.parameters() if p.requires_grad], + lr=cfg.lr_dedicated, + ) + BX._save_checkpoint(policy, os.path.join(outdir, "round_00.pt"), dict( + round=0, experiment=cfg.arm_name, source_sha256=checkpoint_sha, + recipe=asdict(cfg), + )) + preflight = [ + BX.Replica(s, g, n_ped=environment["n_ped"], + ped_speed_range=tuple(environment["ped_speed_range"])) + for s in SP.expansion_scenarios(1, smoke=cfg.smoke) + for g in SP.GAMMAS + ] + ell0, ell, _ = OE._initial_lengthscale(policy, preflight, cfg, device) + history = [] + previous_shard = None + eval_cache = os.path.join(outdir, "m10_cache") + with ProcessPoolExecutor(max_workers=cfg.verifier_workers) as executor: + for round_i in range(1, cfg.rounds + 1): + start = time.perf_counter() + replicas = [ + BX.Replica(s, g, n_ped=environment["n_ped"], + ped_speed_range=tuple( + environment["ped_speed_range"])) + for s in SP.expansion_scenarios(round_i, smoke=cfg.smoke) + for g in SP.GAMMAS + ] + policy.eval() + phi_policy = copy.deepcopy(policy).eval() + for parameter in phi_policy.parameters(): + parameter.requires_grad_(False) + gp, gp_ids, gp_selection = OE.gp_from_previous( + phi_policy, previous_shard, round_i=round_i, ell=ell, + cap=OE.CAP, lam=cfg.gp_lam, phi_s=cfg.phi_s, device=device, + seed=cfg.seed + round_i * 101, + ) + beta, ess = OE._calibrate_beta( + phi_policy, gp, replicas, cfg, device, round_i=round_i, + ) + shard = OS.ExecutedRoundShard(round_i) + gather = OE.gather_offline_round( + policy, phi_policy, gp, beta, replicas, cfg, shard, device, + executor, round_i=round_i, + ) + shard.save(os.path.join( + outdir, "round_shards", f"round_{round_i:02d}.pt", + )) + replay = OR.replay( + policy, optimizer, shard, alpha=cfg.alpha, + exposure_epochs=cfg.exposure_epochs, batch=cfg.batch, + device=device, seed=cfg.seed + round_i * 1_000_003, + ) + pre_path = os.path.join( + outdir, f"round_{round_i:02d}_pre_block.pt", + ) + BX._save_checkpoint(policy, pre_path, dict( + round=round_i, phase="post_ordinary_pre_block", + experiment=cfg.arm_name, recipe=asdict(cfg), + )) + + # ---- D_MPC+ harvest (separate buffer; never enters D/GP) ---- + pop_a, pop_b, pop_stats = AUG.tag_populations(shard) + hard = {int(w["window_id"]): w for w in pop_b} + for window in shard.windows: + if window.get("nvp_context"): + hard[int(window["window_id"])] = window + policy.eval() + records, harvest_audit = MP.harvest_round( + policy, shard, list(hard.values()), executor, + device=device, environment=environment, + max_contexts=MAX_HARVEST_CONTEXTS, + ) + torch.save( + dict(round=round_i, records=records, audit=harvest_audit), + os.path.join(outdir, f"d_mpc_plus_round_{round_i:02d}.pt"), + ) + + audit_before = dict( + local=local_mpc_audit( + policy, shard, records, device=device, + executor=executor, seed_tag="fixed", + ), + m10=m10_eval( + pre_path, f"r{2 * round_i - 1}", + outdir=os.path.join( + outdir, "m10", f"round_{round_i:02d}_pre", + ), + cache_dir=eval_cache, workers=cfg.verifier_workers, + device=device, + ), + ) + block = MP.distill_block( + policy, optimizer_dedicated, shard, records, + epochs=cfg.epochs_dedicated, batch=cfg.batch, + seed=cfg.seed + round_i * 7_000_003, + ) + if BS.module_sha256(policy.enc_grid) != encoder_sha: + raise RuntimeError("visual encoder changed") + post_path = os.path.join(outdir, f"round_{round_i:02d}.pt") + BX._save_checkpoint(policy, post_path, dict( + round=round_i, phase="post_block", + experiment=cfg.arm_name, recipe=asdict(cfg), + )) + audit_after = dict( + local=local_mpc_audit( + policy, shard, records, device=device, + executor=executor, seed_tag="fixed", + ), + m10=m10_eval( + post_path, f"r{2 * round_i}", + outdir=os.path.join( + outdir, "m10", f"round_{round_i:02d}_post", + ), + cache_dir=eval_cache, workers=cfg.verifier_workers, + device=device, + ), + ) + record = dict( + round=round_i, experiment=cfg.arm_name, + beta=float(beta), calibrated_ess=float(ess), + gather_counts=gather["counts"], + outcomes=gather["outcomes"], + replay=dict( + optimizer_steps=replay["optimizer_steps"], + positive=replay["positive_eligible"], + negative=replay["negative_eligible"], + ), + populations=pop_stats, + harvest=dict( + counts=harvest_audit["counts"], + per_gamma=harvest_audit["per_gamma"], + ), + distill_block=block, + audit_before=audit_before, + audit_after=audit_after, + checkpoints=dict(pre=pre_path, post=post_path), + wall_seconds=time.perf_counter() - start, + ) + history.append(record) + with open(os.path.join(outdir, "metrics.jsonl"), "a") as stream: + stream.write(json.dumps(record, allow_nan=False) + "\n") + print(json.dumps(dict( + round=round_i, arm=cfg.arm_name, + kept=harvest_audit["counts"]["kept"], + block_steps=block["steps"], + m10_CR_before=audit_before["m10"]["CR"], + m10_CR_after=audit_after["m10"]["CR"], + m10_V_before=audit_before["m10"]["Validity"], + m10_V_after=audit_after["m10"]["Validity"], + socp_rate_before=audit_before["local"].get( + "socp_positive_rate"), + socp_rate_after=audit_after["local"].get( + "socp_positive_rate"), + wall=record["wall_seconds"], + )), flush=True) + previous_shard = shard + + OE._write_json(os.path.join(outdir, "COMPLETE.json"), dict( + status="CLAUDE_MPC_DISTILL_COMPLETE", + experiment=cfg.arm_name, recipe=asdict(cfg), + source_checkpoint_sha256=checkpoint_sha, + environment=environment, + m10_bank=dict(ep0=M10_EP0, noise_seed=M10_NOISE_SEED, m_per_gamma=10), + constants=dict( + ell=ell, ell0=ell0, keep_per_context=MP.KEEP_PER_CONTEXT, + max_harvest_contexts=MAX_HARVEST_CONTEXTS, + separation=( + "D_MPC+ is per-round, never enters D/D+/GP/acquisition; " + "privileged controller never used at evaluation" + ), + ), + history=history, + )) + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--checkpoint", required=True) + parser.add_argument("--outdir", required=True) + parser.add_argument("--lr-dedicated", type=float, required=True) + parser.add_argument("--epochs-dedicated", type=int, required=True) + parser.add_argument("--rounds", type=int, default=4) + parser.add_argument("--verifier-workers", type=int, default=8) + parser.add_argument("--device", default="cuda:0") + parser.add_argument("--smoke", action="store_true") + parser.add_argument("--tag", default="mpc") + args = parser.parse_args(argv) + cfg = MPCConfig( + lr_dedicated=args.lr_dedicated, + epochs_dedicated=args.epochs_dedicated, + rounds=args.rounds, verifier_workers=args.verifier_workers, + smoke=args.smoke, tag=args.tag, + ) + run(args.checkpoint, args.outdir, cfg, device=args.device) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/claude_mpc_pool.py b/overnight_run_07_12_sfm/claude_mpc_pool.py new file mode 100644 index 0000000..17fa60c --- /dev/null +++ b/overnight_run_07_12_sfm/claude_mpc_pool.py @@ -0,0 +1,419 @@ +"""Privileged MPC candidate-pool harvesting into a separate D_MPC+ buffer. + +Imports the original Codex privileged candidate-pool logic (read-only source: +``agent/sfm-adhoc-controller-overlay-20260726`` @ 0e0eca2, +``sfm_adhoc_controller_compare.privileged_sfm_config`` and the pool machinery +in ``sfm_kazuki``) without modifying either. At declared hard contexts of an +ordinary B1 gathering round this module: + +1. reconstructs the LIVE reactive SFM crowd by deterministic prefix replay of + the episode's executed actions (verified against the stored context); +2. regenerates the original Codex MPC pool at that context: guided flow + sampling (n_sample=200) + MPPI refinement + brake/goal/avoidance templates + + 25 constant-acceleration escapes, exactly as + ``exact_sfm_horizon_filter_action`` constructs it; +3. labels every candidate with BOTH the privileged exact-SFM look-ahead + feasibility (recoverable AND horizon clearance >= the per-gamma hard + margin) and our canonical exact full-H10 SOCP verifier; +4. keeps only the intersection, ranks it by the native frozen SafeMPPI + proposal cost, and stores the top ``KEEP_PER_CONTEXT`` (context, U) pairs + in a fresh per-round ``D_MPC+`` buffer. + +``D_MPC+`` NEVER enters the ordinary D/D+, the GP buffer, or acquisition. +Records carry no x0: distillation uses ``policy.cfm_loss`` which draws fresh +Gaussian CFM bases at every step. The privileged controller itself is never +used at evaluation time. +""" +from __future__ import annotations + +import numpy as np +import torch + +import _paths # noqa: F401 +import sfm_b1_cost as BC +import sfm_b1_store as BS +import sfm_kazuki as KZ +import sfm_metrics2 as SM +import sfm_scene as SS + +KEEP_PER_CONTEXT = 2 +POOL_SEED = 700_000 +H = 10 + + +def privileged_sfm_config(): + """Verbatim import of the historical v3 wrapper recipe. + + Source: agent/sfm-adhoc-controller-overlay-20260726 @ 0e0eca2, + sfm_adhoc_controller_compare.privileged_sfm_config (read-only). + """ + gammas = tuple(map(float, SS.GAMMAS)) + return KZ.KazukiConfig( + safe_coefs=(0.3,), + goal_coef=0.5, + n_sample=200, + n_elite=10, + n_copy=200, + exact_sfm_step_filter=True, + step_filter_margin=0.22, + step_filter_horizon=10, + step_filter_goal_plans=12, + step_filter_avoid_plans=18, + step_filter_always_select=True, + step_filter_min_progress=0.05, + step_filter_goal_score_weight=1.0, + step_filter_clearance_weight=0.05, + step_filter_escape_patience=5, + step_filter_escape_burst=3, + step_filter_viability_lookahead=20, + step_filter_viability_band=0.05, + step_filter_viability_escalate=True, + step_filter_viability_escalation_band=1.0, + step_filter_viability_escalation_min_progress=2.0, + step_filter_viability_escalation_entry_progress=5.0, + step_filter_viability_escalation_burst=40, + step_filter_stagnation_gamma_max=0.1, + step_filter_stagnation_window=20, + step_filter_stagnation_progress=0.1, + step_filter_stagnation_horizon=20, + step_filter_stagnation_burst=4, + controller_gammas=gammas, + safe_coef_by_gamma=(1.0, 0.3, 1.0, 0.3, 0.3, 0.3, 0.1), + goal_coef_by_gamma=(2.0, 0.5, 2.0, 0.5, 0.5, 0.5, 3.0), + step_filter_margin_by_gamma=(0.24, 0.22, 0.24, 0.22, 0.22, 0.22, 0.22), + step_filter_goal_score_weight_by_gamma=(2.0, 1.0, 2.0, 1.0, 1.0, 6.0, 2.0), + step_filter_clearance_weight_by_gamma=(0.05, 0.05, 0.05, 0.05, 0.05, 0.05, 0.0), + step_filter_clearance_target_weight_by_gamma=(0.0,) * len(gammas), + ).validate() + + +def replay_prefix_humans(shard, scenario, gamma, upto_step, environment): + """Deterministically reconstruct the live SFM crowd at a stored context.""" + windows_by_step = {} + for window in shard.windows: + context = shard.contexts[int(window["context_id"])] + if ( + int(context["scenario_id"]) == int(scenario) + and round(float(context["gamma"]), 8) == round(float(gamma), 8) + ): + windows_by_step[int(context["step"])] = window + humans = SS.make_humans( + int(scenario), 0, int(environment["n_ped"]), + tuple(environment["ped_speed_range"]), + ) + state = np.zeros(4, np.float32) + for step in range(int(upto_step)): + window = windows_by_step.get(step) + if window is None: + raise KeyError( + f"missing executed window for s{scenario} g{gamma} step {step}" + ) + action = np.asarray(window["controls"], np.float32)[0] + state[:2] = state[:2] + SS.DT * state[2:4] + 0.5 * SS.DT ** 2 * action + state[2:4] = state[2:4] + SS.DT * action + SS.advance_humans(humans, state) + return humans, state + + +def _recoverable(inside, terminal, reach_step, horizon): + stop_hi = terminal[:, :2] + np.maximum(terminal[:, 2:4], 0.0) ** 2 / (2.0 * SS.U_MAX) + stop_lo = terminal[:, :2] - np.maximum(-terminal[:, 2:4], 0.0) ** 2 / (2.0 * SS.U_MAX) + reached = reach_step <= horizon + return inside & ( + reached + | ((stop_hi <= SS.TASK_HI).all(axis=1) & (stop_lo >= SS.TASK_LO).all(axis=1)) + ) + + +def _deduplicate_plans_with_sources(plans): + """Extend, bound, and de-duplicate plans while preserving first provenance.""" + unique, sources = [], [] + for plan, source in plans: + plan = KZ._extend_plan_with_goal( + np.asarray(source["state"], np.float32), plan, H, + ) + plan = np.clip(np.asarray(plan, np.float32)[:H], -SS.U_MAX, SS.U_MAX) + if not any(np.allclose(plan, old, atol=1e-7) for old in unique): + unique.append(plan) + sources.append({ + key: value for key, value in source.items() if key != "state" + }) + return unique, sources + + +def build_codex_pool( + policy, context, humans, *, device, seed_step, track_sources=False, +): + """Regenerate the original Codex MPC pool at one stored context. + + Deliberately NOT wrapped in ``torch.no_grad``: the Codex guidance + differentiates its CBF/goal rewards with respect to the latent inside + ``guided_generate`` (policy weights are only evaluated under its own + internal ``no_grad`` and are never modified here). + """ + gamma = float(context["gamma"]) + base = privileged_sfm_config() + cfg = KZ._gamma_controller_config(base, gamma).validate() + guidance_cfg = KZ._gamma_guidance_config(cfg, gamma) + state = np.asarray(context["state"], np.float32) + ped_xy = np.asarray(context["ped_xy"], np.float32) + ped_vel = np.asarray(context["ped_vel"], np.float32) + hp10 = torch.as_tensor(context["hp10"], device=device)[None].float() + low = torch.as_tensor(context["low5"], device=device)[None].float() + hist = torch.as_tensor(context["hist"], device=device)[None].float() + ctx = policy.ctx_from(hp10, low, hist).squeeze(0) + goal = torch.tensor(SS.GOAL, dtype=torch.float32, device=device) + torch.manual_seed( + POOL_SEED + int(context["scenario_id"]) * 1000 + int(seed_step) + ) + z = torch.randn(int(cfg.n_sample), int(policy.d), device=device) + taus = cfg.ode_times + ped_pred = KZ.predict_pedestrians_t(ped_xy, ped_vel, H, SS.DT, device) + ped_vel_t = torch.tensor(ped_vel, dtype=torch.float32, device=device) + z1, _, _ = KZ.guided_generate( + policy, ctx, state, goal, ped_pred, ped_vel_t, + SS.R_PED + cfg.collision_margin, z, taus, guidance_cfg, + collect_diagnostics=False, + ) + u_gen = torch.clamp( + z1.reshape(int(cfg.n_sample), H, 2) * float(policy.u_max), + -float(policy.u_max), float(policy.u_max), + ) + u_best, refine_diag = KZ.flow_mppi_refine( + policy, state, goal, ped_pred, SS.R_PED + cfg.collision_margin, + u_gen, None, guidance_cfg, collect_diagnostics=True, + ) + refined_pool = refine_diag.pop("_refined_controls") + nominal = u_best.detach().cpu().numpy().astype(np.float32) + plans = [( + nominal, + dict(family="nominal", source_index=0, state=state), + )] + plans.extend(( + plan, + dict(family="refined", source_index=index, state=state), + ) for index, plan in enumerate(np.asarray(refined_pool, np.float32))) + plans.append(( + KZ._brake_control_plan(state, H), + dict(family="brake", source_index=0, state=state), + )) + plans.extend(( + plan, + dict(family="goal", source_index=index, state=state), + ) for index, plan in enumerate(KZ._goal_control_plans( + state, int(cfg.step_filter_goal_plans), H, + ))) + plans.extend(( + plan, + dict(family="avoidance", source_index=index, state=state), + ) for index, plan in enumerate(KZ._avoidance_control_plans( + humans, state, int(cfg.step_filter_avoid_plans), H, + ))) + const_index = 0 + for ux in np.linspace(-SS.U_MAX, SS.U_MAX, 5): + for uy in np.linspace(-SS.U_MAX, SS.U_MAX, 5): + plans.append(( + np.repeat(np.array([[ux, uy]], np.float32), H, axis=0), + dict( + family="constant_acceleration", + source_index=const_index, state=state, + ), + )) + const_index += 1 + unique, sources = _deduplicate_plans_with_sources(plans) + stacked = np.stack(unique) + clear, inside, terminal, _, reach_step = KZ._simulate_sfm_plans( + humans, state, stacked, H, + ) + margin = KZ._adaptive_step_filter_margin(cfg, gamma) + privileged = _recoverable(inside, terminal, reach_step, H) & ( + clear >= float(margin) + ) + result = dict( + plans=stacked, + privileged_feasible=privileged, + privileged_clearance=clear, + margin=float(margin), + pool_manifest=dict( + nominal=1, refined=len(refined_pool), + brake=1, goal_plans=int(cfg.step_filter_goal_plans), + avoid_plans=int(cfg.step_filter_avoid_plans), + const_accel=25, unique=len(unique), + ), + ) + if track_sources: + result["candidate_sources"] = sources + result["nominal_plan"] = np.asarray(nominal, np.float32).copy() + result["refined_pool"] = np.asarray(refined_pool, np.float32).copy() + return result + + +def harvest_round( + policy, shard, hard_windows, executor, *, device, environment, + max_contexts=None, +): + """Build the per-round D_MPC+ from the round's declared hard contexts.""" + by_lineage = {} + for window in hard_windows: + context = shard.contexts[int(window["context_id"])] + key = (int(context["scenario_id"]), round(float(context["gamma"]), 8)) + by_lineage.setdefault(key, []).append(int(window["context_id"])) + records, audit_rows = [], [] + counts = dict( + hard_contexts=0, replay_mismatch=0, pool_candidates=0, + privileged_feasible=0, socp_positive=0, intersection=0, kept=0, + ) + context_ids_all = sorted( + cid for ids in by_lineage.values() for cid in ids + ) + if max_contexts is not None and len(context_ids_all) > int(max_contexts): + stride = len(context_ids_all) / float(max_contexts) + context_ids_all = [ + context_ids_all[int(i * stride)] for i in range(int(max_contexts)) + ] + chosen = set(context_ids_all) + for (scenario, gamma), context_ids in sorted(by_lineage.items()): + for context_id in sorted(context_ids): + if context_id not in chosen: + continue + context = shard.contexts[context_id] + counts["hard_contexts"] += 1 + humans, replay_state = replay_prefix_humans( + shard, scenario, gamma, int(context["step"]), environment, + ) + ped_xy_live, _ = SS.collect_humans(humans) + if not np.allclose( + ped_xy_live, np.asarray(context["ped_xy"], np.float32), + atol=1e-4, + ) or not np.allclose( + replay_state, np.asarray(context["state"], np.float32), + atol=1e-4, + ): + counts["replay_mismatch"] += 1 + continue + pool = build_codex_pool( + policy, context, humans, device=device, + seed_step=int(context["step"]), + ) + plans = pool["plans"] + counts["pool_candidates"] += len(plans) + privileged = pool["privileged_feasible"] + counts["privileged_feasible"] += int(privileged.sum()) + tasks = [ + (index, 0, context["state"], plans[index], + context["ped_xy"], context["ped_vel"], gamma) + for index in range(len(plans)) if privileged[index] + ] + results = {i: r for i, _, r in executor.map( + SM.verify_in_worker, tasks, + )} + certified = [ + index for index in results + if results[index].get("resolved") + and int(results[index].get("y", 0)) == 1 + ] + counts["socp_positive"] += len(certified) + counts["intersection"] += len(certified) + if not certified: + continue + controls = torch.as_tensor( + np.stack([plans[i] for i in certified]), dtype=torch.float32, + ) + costs = BC.safemppi_proposal_cost( + context["state"], controls, SS.GOAL, + context["ped_xy"], context["ped_vel"], + ).cpu().numpy() + order = sorted( + range(len(certified)), key=lambda j: (float(costs[j]), j), + ) + for rank, j in enumerate(order[:KEEP_PER_CONTEXT]): + index = certified[j] + counts["kept"] += 1 + records.append(dict( + context_id=int(context_id), + controls=np.asarray(plans[index], np.float32), + y=1, + query_id=len(records), + source="codex_privileged_mpc_pool", + privileged_clearance=float( + pool["privileged_clearance"][index], + ), + privileged_margin=pool["margin"], + safemppi_cost=float(costs[j]), + rank=int(rank), + verifier_diagnostics=dict( + results[index]["diagnostics"], + ), + )) + audit_rows.append(dict( + context_id=int(context_id), scenario=int(scenario), + gamma=float(gamma), step=int(context["step"]), + rank=int(rank), cost=float(costs[j]), + privileged_clearance=float( + pool["privileged_clearance"][index], + ), + socp_slack=float( + results[index]["diagnostics"]["slack"], + ), + )) + per_gamma = {} + for record in records: + gamma = str(shard.contexts[record["context_id"]]["gamma"]) + per_gamma[gamma] = per_gamma.get(gamma, 0) + 1 + return records, dict(counts=counts, per_gamma=per_gamma, rows=audit_rows) + + +class MPCView: + """Duck-typed positive-only view over D_MPC+ for the mass machinery.""" + + def __init__(self, shard, records): + self.round_i = int(shard.round_i) + self.contexts = shard.contexts + self.windows = [dict(record) for record in records] + for index, row in enumerate(self.windows): + row["window_id"] = index + row["query_id"] = index + + @property + def Dplus(self): + return list(self.windows) + + +def distill_block(policy, optimizer, shard, records, *, epochs, batch, seed): + """Dedicated CFM distillation on D_MPC+ only (fresh Gaussian bases).""" + if not records: + return dict(steps=0, records=0, losses=None) + view = MPCView(shard, records) + pairs = [(view, row) for row in view.windows] + mass, accounting = BS.hierarchy_mass(pairs) + policy.train() + losses = [] + steps = 0 + generator = np.random.default_rng(int(seed)) + device = next(policy.parameters()).device + for epoch in range(int(epochs)): + order = list(generator.permutation(len(pairs))) + for start in range(0, len(order), int(batch)): + chunk = [pairs[i] for i in order[start:start + int(batch)]] + grid, low, hist, controls = BS._tensor_batch(chunk, device) + context = policy.ctx_from(grid, low, hist) + weights = torch.as_tensor([ + len(pairs) * mass[(id(view), int(row["query_id"]))] + for _, row in chunk + ], dtype=controls.dtype, device=device) + torch.manual_seed(int(seed) + epoch * 100_003 + start) + loss = policy.cfm_loss(controls, context, weights=weights) + if not bool(torch.isfinite(loss)): + raise FloatingPointError("non-finite MPC distillation loss") + optimizer.zero_grad(set_to_none=True) + loss.backward() + optimizer.step() + losses.append(float(loss.detach())) + steps += 1 + policy.eval() + return dict( + steps=steps, records=len(pairs), epochs=int(epochs), + loss_first=losses[0], loss_last=losses[-1], + loss_mean=float(np.mean(losses)), + mass_gamma=accounting["gamma"], + ) diff --git a/overnight_run_07_12_sfm/claude_mpc_select.py b/overnight_run_07_12_sfm/claude_mpc_select.py new file mode 100644 index 0000000..b691a5a --- /dev/null +++ b/overnight_run_07_12_sfm/claude_mpc_select.py @@ -0,0 +1,130 @@ +"""Apply the pre-registered MPC-study M10 selection rule. + +Collects every saved checkpoint cell (pre-block and post-block, every arm and +round) from the sweep's per-arm ``m10/`` audit directories, plus the common r0 +and the no-distillation control cells, and applies the rule declared in +MPC_STUDY_PREREGISTRATION.json verbatim: + + liveness gate: SR >= SR(r0) - 0.02 and timeout <= timeout(r0) + 0.05; + among eligible: min CR, then max Validity, then max successful clearance, + then min successful time-to-goal; + the winner must be a POST-BLOCK checkpoint to count as a distillation + effect; its pre-block sibling is always reported alongside. +""" +from __future__ import annotations + +import argparse +import glob +import json +import os + + +def _pooled(path): + with open(path) as stream: + payload = json.load(stream) + out = [] + for record in payload["records"]: + cell = record["cell"] + p = cell["summary"]["pooled"] + out.append(dict( + label=record["label"], + checkpoint=cell["checkpoint"], + checkpoint_sha256=cell["checkpoint_sha256"], + SR=float(p["SR"]), CR=float(p["CR"]), + timeout=float(p["timeout"]), + Validity=float(p["Validity"]["mean"]), + clearance=p["successful_clearance"]["mean"], + time=p["successful_time_to_goal"]["mean"], + )) + return out + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--sweep-root", required=True) + parser.add_argument("--arms", nargs="+", required=True) + parser.add_argument("--r0-control-metrics", required=True) + parser.add_argument("--out", required=True) + args = parser.parse_args(argv) + + baseline = _pooled(args.r0_control_metrics) + r0 = next(row for row in baseline if row["label"] == "r0") + control_rows = [ + dict(row, arm="control_no_distill", phase="control", + round=int(row["label"][1:])) + for row in baseline if row["label"] != "r0" + ] + cells = [] + for arm in args.arms: + for path in sorted(glob.glob(os.path.join( + args.sweep_root, arm, "m10", "round_*_*", + "raw_m10_offline_metrics.json", + ))): + phase = "post" if path.split(os.sep)[-2].endswith("_post") else "pre" + round_i = int(path.split(os.sep)[-2].split("_")[1]) + for row in _pooled(path): + cells.append(dict( + row, arm=arm, phase=phase, round=round_i, + )) + gate_sr = r0["SR"] - 0.02 + gate_to = r0["timeout"] + 0.05 + eligible = [ + c for c in cells + if c["SR"] >= gate_sr and c["timeout"] <= gate_to + and c["clearance"] is not None and c["time"] is not None + ] + + def key(c): + return ( + c["CR"], -c["Validity"], + -(c["clearance"] if c["clearance"] is not None else -1), + c["time"] if c["time"] is not None else 1e9, + c["round"], c["arm"], + ) + + ordered = sorted(eligible, key=key) + post_ordered = [c for c in ordered if c["phase"] == "post"] + winner = post_ordered[0] if post_ordered else None + sibling = None + if winner: + sibling = next( + (c for c in cells + if c["arm"] == winner["arm"] and c["round"] == winner["round"] + and c["phase"] == "pre"), + None, + ) + payload = dict( + status="MPC_M10_SELECTION_APPLIED", + rule="MPC_STUDY_PREREGISTRATION.json verbatim", + bank=dict(ep0=350000, noise_seed=20260733, m_per_gamma=10), + r0=r0, + liveness_gate=dict(SR_min=gate_sr, timeout_max=gate_to), + n_cells=len(cells), n_eligible=len(eligible), + winner_post_block=winner, + pre_block_sibling=sibling, + top_eligible=ordered[:8], + control_no_distill=control_rows, + all_cells=sorted( + cells + control_rows, + key=lambda c: (c["arm"], c["round"], c.get("phase", "")), + ), + ) + with open(args.out, "w") as stream: + json.dump(payload, stream, indent=1, allow_nan=False) + print(json.dumps(dict( + r0={k: r0[k] for k in ("SR", "CR", "Validity")}, + winner=None if winner is None else { + k: winner[k] + for k in ("arm", "round", "phase", "SR", "CR", "Validity", + "clearance", "time") + }, + sibling=None if sibling is None else { + k: sibling[k] + for k in ("SR", "CR", "Validity", "clearance", "time") + }, + eligible=len(eligible), + ), indent=1)) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/claude_offline_aug.py b/overnight_run_07_12_sfm/claude_offline_aug.py new file mode 100644 index 0000000..d46f874 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_offline_aug.py @@ -0,0 +1,507 @@ +"""Opt-in replay-data interventions for offline executed-window SFM expansion. + +Everything in this module is additive and OFF by default: with +``replay_mode="original"`` the caller uses ``sfm_b1_offline_replay.replay`` on +the untouched :class:`ExecutedRoundShard` and no code here runs. The exact +B1 gathering, GP acquisition, B querying, verifier, and raw evaluation are +never modified here; this module only re-composes which resolved executed +records (plus optional exact-certified synthetic recovery records) the replay +step trains on. + +Declared population rules (fixed BEFORE evaluating their effect): + +Population A - near-obstacle verified positives. An executed window with +``y=1`` whose plan-time constant-velocity predicted minimum time-indexed +pedestrian clearance is at most ``NEAR_CLEARANCE_MAX`` metres and whose +10-step displacement is at least ``MIN_DISPLACEMENT`` metres (successful +avoidance = certified, close to a constraint boundary, and actually moving; +the displacement floor equals the existing declared trap displacement). + +Population B - collision/trap/no-progress negatives. An executed window with +``y=0`` satisfying at least one of: actual collision after the executed +action (``collision_after_action``); plan-time predicted window collision +(``collision_free`` false); the existing declared trap rule fired at this +step (``trap_event``: closed-loop displacement < 0.2 m over the last 10 +executed steps); or window displacement below ``MIN_DISPLACEMENT`` (the same +declared 10-step/0.2 m insufficient-progress rule applied to the executed +window). A certified-but-slow window (y=1) is never relabeled negative. + +Population C - deterministic certified recovery positives. For each hard +context (the context of a population-B window), solve the deterministic +control problem written out in :func:`recovery_candidates`: + + minimize J(u) = || p_10(u) - GOAL ||_2 + over the fixed finite deterministic family below + subject to |u_t|_inf <= U_MAX, double-integrator dynamics with DT, + and the EXACT full-H10 verifier certificate y=1 + (task-space bounds, time-indexed CV collision avoidance, + GREEN moving-window SOCP certificate). + +Family: closed-loop brake steps b in {0,2,4} (u_t = clip(-v_t/DT)) followed +by bang-bang steering u = a*(cos t_k, sin t_k) for the first half of the +remaining steps and -a*(...) for the second half, over K_DIR=16 world +directions and magnitudes a in {0.7, 1.4, 2.0}; plus the pure 10-step brake. +Candidates are pre-filtered by the cheap exact numpy task-space and +CV-collision checks, ranked by J, and at most ``PREVERIFY_CAP`` are submitted +to the exact SOCP verifier; the first ``RECOVERY_KEEP`` certified candidates +(in increasing J) enter the positive set. A failed candidate is discarded, +never relabeled. Recovery rows join the positive replay population at their +parent context, so the existing hierarchical mass (gamma -> episode -> +context -> query) automatically splits the parent context's mass across +them; they carry ``x0 = zeros(20)`` for schema compatibility and are NEVER +eligible for the GP buffer (the GP reads only the real ExecutedRoundShard). +The recovery controller itself is never deployed at evaluation time. +""" +from __future__ import annotations + +import math + +import numpy as np + +import _paths # noqa: F401 +import sfm_b1_offline_replay as OR +import sfm_b1_offline_store as OS +import sfm_metrics2 as SM +import sfm_scene as SS + +NEAR_CLEARANCE_MAX = 0.35 +MIN_DISPLACEMENT = 0.2 +BRAKE_STEPS = (0, 2, 4) +K_DIR = 16 +ACCELS = (0.7, 1.4, 2.0) +PREVERIFY_CAP = 24 +RECOVERY_KEEP = 2 +REPLAY_MODES = ( + "original", "hard", "hard_recovery", "orig_plus_recovery", + "orig_plus_recovery_v2", +) + +# --- Recovery family v2 ("dodge-then-cruise"), declared 2026-07-26 before +# any evaluation of its effect. Motivation (user hypothesis + Stage-A +# measurement): the v1 brake/bang-bang family certifies CONSERVATIVE escapes +# that end near-stationary; training on them creates slowdown. v2 candidates +# end moving TOWARD the goal at cruise speed: +# phase 1 (d in {0,2,3} steps): dodge with u = a*(cos t_k, sin t_k), +# a in {1.4, 2.0}, t_k over K_DIR world directions (d=0 skips the dodge); +# phase 2 (remaining steps): deterministic saturated velocity servo +# u_t = clip(KP_CRUISE * (v_des(p_t) - v_t), +/-U_MAX), +# v_des(p) = v_c * unit(GOAL - p), v_c in {1.0, 1.5}. +# Objective (v2): J2 = ||p_10 - GOAL|| - 0.5 * (v_10 . unit(GOAL - p_10)) — +# prefer end states that are close to AND moving toward the goal. The same +# cheap exact prefilter, PREVERIFY_CAP, RECOVERY_KEEP, and the exact full-H10 +# SOCP certificate gate apply unchanged. +DODGE_STEPS = (0, 2, 3) +CRUISE_SPEEDS = (1.0, 1.5) +KP_CRUISE = 4.0 +V2_ACCELS = (1.4, 2.0) + + +def declared_rules(): + return dict( + population_A=dict( + requires="y=1", + predicted_min_clearance_max=NEAR_CLEARANCE_MAX, + window_displacement_min=MIN_DISPLACEMENT, + ), + population_B=dict( + requires="y=0", + any_of=[ + "collision_after_action", + "not collision_free (plan-time predicted window collision)", + "trap_event (declared closed-loop 10-step/0.2m rule)", + f"window displacement < {MIN_DISPLACEMENT} m over H=10", + ], + ), + population_C=dict( + objective="min ||p_10(u) - GOAL||_2 over the fixed family", + family=dict( + brake_steps=list(BRAKE_STEPS), k_dir=K_DIR, + accels=list(ACCELS), plus="pure 10-step closed-loop brake", + ), + preverify_cap=PREVERIFY_CAP, + keep_per_context=RECOVERY_KEEP, + certificate="exact full-H10 SM.verify_query y=1 only", + ), + ) + + +def _window_geometry(context, controls): + segment = SM.rollout_positions(context["state"], controls) + prediction = SM.predict_pedestrians( + context["ped_xy"], context["ped_vel"], H=len(controls), + ) + clearance = float( + np.linalg.norm(segment[:, None, :] - prediction, axis=2).min() + - SS.R_PED + ) + displacement = float(np.linalg.norm(segment[-1] - segment[0])) + return clearance, displacement + + +def tag_populations(shard): + """Classify every executed window; returns (popA, popB, stats).""" + pop_a, pop_b = [], [] + reasons = dict(actual_collision=0, predicted_collision=0, trap=0, + no_progress=0) + for window in shard.windows: + context = shard.contexts[int(window["context_id"])] + clearance, displacement = _window_geometry( + context, window["controls"], + ) + if int(window["y"]) == 1: + if ( + clearance <= NEAR_CLEARANCE_MAX + and displacement >= MIN_DISPLACEMENT + ): + pop_a.append(window) + else: + actual = bool(window.get("collision_after_action")) + predicted = not bool(window["collision_free"]) + trap = bool(window.get("trap_event")) + slow = displacement < MIN_DISPLACEMENT + if actual or predicted or trap or slow: + pop_b.append(window) + reasons["actual_collision"] += int(actual) + reasons["predicted_collision"] += int(predicted) + reasons["trap"] += int(trap) + reasons["no_progress"] += int(slow) + stats = dict( + D=len(shard.windows), Dplus=len(shard.Dplus), + Dminus=len(shard.Dminus), popA=len(pop_a), popB=len(pop_b), + popB_reasons=reasons, rules=declared_rules(), + ) + return pop_a, pop_b, stats + + +def _closed_loop_brake(state, steps): + """Deterministic max-effort brake controls for ``steps`` steps.""" + velocity = np.asarray(state, np.float32).reshape(4)[2:4].copy() + controls = [] + for _ in range(steps): + action = np.clip(-velocity / SS.DT, -SS.U_MAX, SS.U_MAX) + controls.append(action.astype(np.float32)) + velocity = velocity + SS.DT * action + return controls, velocity + + +def recovery_candidates(state): + """The fixed deterministic candidate family for one context.""" + candidates = [] + full_brake, _ = _closed_loop_brake(state, 10) + candidates.append(( + np.asarray(full_brake, np.float32), + dict(kind="brake10", brake=10, theta=None, accel=None), + )) + for brake in BRAKE_STEPS: + prefix, _ = _closed_loop_brake(state, brake) + remaining = 10 - brake + forward = math.ceil(remaining / 2) + for k in range(K_DIR): + theta = 2.0 * math.pi * k / K_DIR + direction = np.array( + [math.cos(theta), math.sin(theta)], np.float32, + ) + for accel in ACCELS: + steer = [accel * direction] * forward + steer += [-accel * direction] * (remaining - forward) + controls = np.asarray(prefix + steer, np.float32) + if controls.shape != (10, 2): + raise AssertionError("recovery candidate must be H=10") + candidates.append(( + np.clip(controls, -SS.U_MAX, SS.U_MAX), + dict(kind="brake_steer", brake=brake, + theta=round(theta, 6), accel=accel), + )) + return candidates + + +def _cruise_controls(position, velocity, steps, v_cruise): + """Deterministic saturated velocity servo toward the goal.""" + position = np.asarray(position, np.float32).copy() + velocity = np.asarray(velocity, np.float32).copy() + controls = [] + for _ in range(steps): + offset = SS.GOAL - position + norm = float(np.linalg.norm(offset)) + v_des = ( + v_cruise * offset / norm if norm > 1e-6 + else np.zeros(2, np.float32) + ) + action = np.clip( + KP_CRUISE * (v_des - velocity), -SS.U_MAX, SS.U_MAX, + ).astype(np.float32) + controls.append(action) + position = position + SS.DT * velocity + 0.5 * SS.DT ** 2 * action + velocity = velocity + SS.DT * action + return controls + + +def recovery_candidates_v2(state, ped_xy=None, ped_vel=None): + """Goal-directed dodge-then-cruise family (state-only, deterministic).""" + del ped_xy, ped_vel # verifier inputs; unused by this state-only family + state = np.asarray(state, np.float32).reshape(4) + candidates = [] + for v_cruise in CRUISE_SPEEDS: + controls = _cruise_controls(state[:2], state[2:4], 10, v_cruise) + candidates.append(( + np.asarray(controls, np.float32), + dict(kind="cruise", dodge=0, theta=None, accel=None, + v_cruise=v_cruise), + )) + for dodge in DODGE_STEPS: + if dodge == 0: + continue + for k in range(K_DIR): + theta = 2.0 * math.pi * k / K_DIR + direction = np.array( + [math.cos(theta), math.sin(theta)], np.float32, + ) + for accel in V2_ACCELS: + prefix = [ + np.clip(accel * direction, -SS.U_MAX, SS.U_MAX) + .astype(np.float32) + ] * dodge + position = np.asarray(state[:2], np.float32).copy() + velocity = np.asarray(state[2:4], np.float32).copy() + for action in prefix: + position = ( + position + SS.DT * velocity + + 0.5 * SS.DT ** 2 * action + ) + velocity = velocity + SS.DT * action + for v_cruise in CRUISE_SPEEDS: + controls = prefix + _cruise_controls( + position, velocity, 10 - dodge, v_cruise, + ) + controls = np.asarray(controls, np.float32) + if controls.shape != (10, 2): + raise AssertionError("v2 candidate must be H=10") + candidates.append(( + controls, + dict(kind="dodge_cruise", dodge=dodge, + theta=round(theta, 6), accel=accel, + v_cruise=v_cruise), + )) + return candidates + + +def _objective_v2(context, controls): + segment = SM.rollout_positions(context["state"], controls) + state = np.asarray(context["state"], np.float32).reshape(4) + velocity = state[2:4].copy() + for action in np.asarray(controls, np.float32): + velocity = velocity + SS.DT * action + offset = SS.GOAL - segment[-1] + norm = float(np.linalg.norm(offset)) + toward = float(velocity @ (offset / norm)) if norm > 1e-6 else 0.0 + return norm - 0.5 * toward + + +def _prefilter(context, controls): + """Cheap exact numpy feasibility check + objective J.""" + segment = SM.rollout_positions(context["state"], controls) + if not SM.taskspace_ok(segment): + return None + prediction = SM.predict_pedestrians( + context["ped_xy"], context["ped_vel"], H=10, + ) + if not SM.collision_free_time_indexed(segment, prediction): + return None + return float(np.linalg.norm(segment[-1] - SS.GOAL)) + + +def build_recovery_records(shard, hard_windows, executor, family="v1"): + """Exact-certified recovery positives for the given hard windows. + + Returns (records, audit). Every returned record passed the exact + full-H10 verifier inside ``executor`` (the same worker pool and + ``SM.verify_in_worker`` entry as B1 queries). ``family`` selects the + declared deterministic candidate family and ranking objective: + v1 = brake/bang-bang, J = final goal distance; + v2 = dodge-then-cruise, J2 = goal distance - 0.5 * toward-goal speed. + """ + if family not in ("v1", "v2"): + raise ValueError("recovery family must be v1 or v2") + context_ids = sorted({int(w["context_id"]) for w in hard_windows}) + tasks, meta = [], [] + per_context_pool = {} + for context_id in context_ids: + context = shard.contexts[context_id] + generator_fn = ( + recovery_candidates if family == "v1" + else recovery_candidates_v2 + ) + scored = [] + for controls, provenance in generator_fn(context["state"]): + feasible = _prefilter(context, controls) + if feasible is None: + continue + objective = ( + feasible if family == "v1" + else _objective_v2(context, controls) + ) + scored.append((objective, controls, provenance)) + scored.sort(key=lambda row: (row[0], str(row[2]))) + pool = scored[:PREVERIFY_CAP] + per_context_pool[context_id] = len(pool) + for rank, (objective, controls, provenance) in enumerate(pool): + tasks.append(( + context_id, rank, context["state"], controls, + context["ped_xy"], context["ped_vel"], context["gamma"], + )) + meta.append((context_id, rank, objective, controls, provenance)) + results = list(executor.map(SM.verify_in_worker, tasks)) + verified = {} + for (context_id, rank, result), (_, _, objective, controls, provenance) \ + in zip(results, meta): + verified.setdefault(int(context_id), []).append( + (int(rank), float(objective), controls, provenance, result), + ) + records, audit_rows = [], [] + certified_total = queried_total = 0 + for context_id in context_ids: + rows = sorted(verified.get(context_id, []), key=lambda r: r[0]) + queried_total += len(rows) + kept = 0 + for rank, objective, controls, provenance, result in rows: + if kept >= RECOVERY_KEEP: + break + if not result.get("resolved"): + continue + if int(result.get("y", 0)) != 1 or not bool(result.get("full_h")): + continue + certified_total += 1 + kept += 1 + records.append(dict( + window_id=None, query_id=None, + context_id=int(context_id), + controls=np.asarray(controls, np.float32), + x0=np.zeros(20, np.float32), + y=1, taskspace=True, collision_free=True, certificate=True, + full_h=True, terminal_step=10, train_eligible=True, + execution_source="synthetic_certified_recovery", + nvp_context=False, candidate_id=None, acquisition_step=None, + sigma=None, hp_margin=None, mode="recovery", + verifier_diagnostics=dict(result["diagnostics"]), + recovery_provenance=dict( + parent_round=int(shard.round_i), + parent_context_id=int(context_id), + family=str(family), + generator=provenance, + objective=float(objective), + prefilter_rank=int(rank), + ), + )) + audit_rows.append(dict( + context_id=int(context_id), rank=int(rank), + generator=provenance, J=float(objective), + verifier=dict( + y=int(result["y"]), taskspace=bool(result["taskspace"]), + collision_free=bool(result["collision_free"]), + certificate=bool(result["certificate"]), + slack=float(result["diagnostics"]["slack"]), + ), + )) + audit = dict( + hard_contexts=len(context_ids), + exact_verifier_queries=queried_total, + certified_kept=len(records), + certified_total_seen=certified_total, + keep_per_context=RECOVERY_KEEP, + preverify_cap=PREVERIFY_CAP, + per_context_preverified_pool_mean=( + float(np.mean(list(per_context_pool.values()))) + if per_context_pool else 0.0 + ), + family=str(family), + rows=audit_rows, + ) + return records, audit + + +class ShardView: + """Duck-typed shard exposing a re-composed training population. + + Shares the parent shard's contexts; windows are the selected subset plus + optional synthetic records, re-indexed with dense window/query ids so the + untouched ``sfm_b1_offline_replay.replay`` machinery (hierarchy mass, + stratified batches, exact-once accounting) applies verbatim. + """ + + def __init__(self, parent, windows): + self.round_i = int(parent.round_i) + self.contexts = parent.contexts + self.windows = [] + for index, window in enumerate(windows): + row = dict(window) + row["window_id"] = index + row["query_id"] = index + self.windows.append(row) + + @property + def D(self): + return list(self.windows) + + @property + def Dplus(self): + return [row for row in self.windows if row["y"] == 1] + + @property + def Dminus(self): + return [row for row in self.windows if row["y"] == 0] + + +def build_replay_view(shard, mode, executor=None): + """Compose the replay population for ``mode``; returns (view, report).""" + if mode not in REPLAY_MODES: + raise ValueError(f"replay mode must be one of {REPLAY_MODES}") + if mode == "original": + return shard, dict(mode=mode, note="untouched ExecutedRoundShard") + pop_a, pop_b, stats = tag_populations(shard) + report = dict(mode=mode, populations=stats) + if mode in ("orig_plus_recovery", "orig_plus_recovery_v2"): + # Declared BEFORE evaluation (Stage-A log 2026-07-26): keep the FULL + # original positive and negative populations and only APPEND the + # exact-certified recovery positives at their parent (hard) contexts. + # The _v2 variant uses the declared dodge-then-cruise family. + if executor is None: + raise ValueError("orig_plus_recovery needs the verifier executor") + recovery, audit = build_recovery_records( + shard, pop_b, executor, + family="v2" if mode.endswith("_v2") else "v1", + ) + report["recovery_audit"] = audit + windows = list(shard.windows) + recovery + else: + windows = list(pop_a) + list(pop_b) + if mode == "hard_recovery": + if executor is None: + raise ValueError("hard_recovery needs the verifier executor") + recovery, audit = build_recovery_records(shard, pop_b, executor) + report["recovery_audit"] = audit + windows = windows + recovery + if not any(int(row["y"]) == 1 for row in windows): + # Fail open to the untouched population rather than training on a + # positive-free set (the replay contract requires positives). + report["fallback"] = "no positives in composed set; using original" + return shard, report + view = ShardView(shard, windows) + report["view"] = dict( + D=len(view.windows), Dplus=len(view.Dplus), Dminus=len(view.Dminus), + ) + return view, report + + +def replay_with_mode( + policy, optimizer, shard, *, mode, alpha, exposure_epochs, batch, + device, seed, executor=None, +): + """Opt-in wrapper: ``original`` delegates verbatim to OR.replay.""" + view, report = build_replay_view(shard, mode, executor=executor) + result = OR.replay( + policy, optimizer, view, + alpha=alpha, exposure_epochs=exposure_epochs, batch=batch, + device=device, seed=seed, + ) + result["replay_intervention"] = report + return result diff --git a/overnight_run_07_12_sfm/claude_offline_exec_ext.py b/overnight_run_07_12_sfm/claude_offline_exec_ext.py new file mode 100644 index 0000000..a6c15f0 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_offline_exec_ext.py @@ -0,0 +1,286 @@ +"""Extended offline executed-window expansion runner (opt-in knobs). + +This is an additive experiment driver. It reuses the immutable core of +``sfm_b1_offline_exec`` verbatim — ``gather_offline_round`` (56 lineages, +K=16, B=4 exact verifier queries, executed-window-only store, NVP +continuation), ``gp_from_previous``, ``_calibrate_beta``, +``_initial_lengthscale`` — and differs ONLY in the declared recipe knobs: + +- ``lr`` (optimizer learning rate; control 1e-4), +- ``ess_target`` (acquisition ESS calibration target; control 0.5), +- ``rounds`` (number of macro-rounds), +- ``replay_mode`` in {original, hard, hard_recovery} + (see ``claude_offline_aug``; ``original`` delegates verbatim to + ``sfm_b1_offline_replay.replay``), +- ``alpha`` / ``exposure_epochs`` restricted to the original replay + contract sets {0, 0.01, 0.1} and {1, 10, 100}. + +With (lr=1e-4, ess_target=0.5, rounds=10, replay_mode="original") the run is +behaviourally identical to ``sfm_b1_offline_exec`` (same keyed seeds, same +calls); this equivalence is asserted by comparing round-1 shard digests +against the archived control run. +""" +from __future__ import annotations + +import argparse +from concurrent.futures import ProcessPoolExecutor +import copy +from dataclasses import asdict, dataclass +import json +import os +import time + +import torch + +import _paths # noqa: F401 +import claude_offline_aug as AUG +import grid_policy_sfm as GPS +import sfm_b1_expand as BX +import sfm_b1_offline_exec as OE +import sfm_b1_offline_store as OS +import sfm_b1_store as BS +import sfm_metrics2 as SM +import sfm_protocol as SP +import sfm_scene as SS + + +@dataclass(frozen=True) +class ExtConfig: + alpha: float + exposure_epochs: int + selector: str = "margin" + rounds: int = 10 + lr: float = 1.0e-4 + ess_target: float = 0.5 + replay_mode: str = "original" + K: int = 16 + B: int = 4 + T: int = 180 + H: int = 10 + batch: int = 128 + nfe: int = 8 + temp: float = 1.0 + phi_s: float = 0.9 + gp_lam: float = OE.GP_LAMBDA + verifier_workers: int = 8 + seed: int = 20260724 + scene_profile: str = OE.SCENE_PROFILE + smoke: bool = False + tag: str = "ext" + + def validate(self): + if self.selector not in OE.EXECUTION_SELECTORS: + raise ValueError(f"selector must be one of {OE.EXECUTION_SELECTORS}") + if float(self.alpha) not in OE.ALPHAS: + raise ValueError(f"alpha must be one of {OE.ALPHAS}") + if int(self.exposure_epochs) not in OE.EXPOSURE_EPOCHS: + raise ValueError( + f"exposure_epochs must be one of {OE.EXPOSURE_EPOCHS}" + ) + if self.replay_mode not in AUG.REPLAY_MODES: + raise ValueError(f"replay_mode must be one of {AUG.REPLAY_MODES}") + if not 0.05 <= float(self.ess_target) <= 0.95: + raise ValueError("ess_target out of the studied range") + if not 0.0 < float(self.lr) <= 1.0e-3: + raise ValueError("lr out of the studied range") + if not 1 <= int(self.rounds) <= 20: + raise ValueError("rounds out of the studied range") + # The immutable scientific core is pinned exactly as in the control. + if ( + int(self.K), int(self.B), int(self.T), int(self.H), + int(self.batch), float(self.gp_lam), float(self.temp), + self.scene_profile, int(self.nfe), float(self.phi_s), + ) != (16, 4, 180, 10, 128, OE.GP_LAMBDA, 1.0, OE.SCENE_PROFILE, 8, 0.9): + raise ValueError("immutable offline executed-window core changed") + if int(self.verifier_workers) < 1: + raise ValueError("verifier_workers must be positive") + return self + + @property + def arm_name(self): + alpha = str(float(self.alpha)).replace(".", "p") + lr = f"{self.lr:.0e}".replace("-", "m") + ess = str(float(self.ess_target)).replace(".", "p") + return ( + f"{self.tag}_{self.selector}_a{alpha}_e{int(self.exposure_epochs):03d}" + f"_lr{lr}_ess{ess}_{self.replay_mode}" + ) + + +def run(checkpoint, outdir, cfg, *, device): + cfg.validate() + checkpoint = os.path.abspath(checkpoint) + outdir = os.path.abspath(outdir) + if not os.path.isfile(checkpoint): + raise FileNotFoundError(checkpoint) + checkpoint_sha = OS.sha256_file(checkpoint) + if checkpoint_sha != OE.EXPECTED_CHECKPOINT_SHA256: + raise ValueError( + f"checkpoint SHA mismatch: expected " + f"{OE.EXPECTED_CHECKPOINT_SHA256}, got {checkpoint_sha}" + ) + if os.path.exists(outdir): + raise FileExistsError(f"refusing to reuse output directory: {outdir}") + os.makedirs(outdir) + environment = SS.scene_profile(cfg.scene_profile) + policy, _ = GPS.load_sfm_policy(checkpoint, device=device) + frozen_parameters = BS.configure_expansion_trainability(policy) + visual_encoder_sha = BS.module_sha256(policy.enc_grid) + optimizer = torch.optim.Adam( + [p for p in policy.parameters() if p.requires_grad], lr=cfg.lr, + ) + BX._save_checkpoint(policy, os.path.join(outdir, "round_00.pt"), dict( + round=0, experiment=cfg.arm_name, source_checkpoint=checkpoint, + source_sha256=checkpoint_sha, encoder_sha256=visual_encoder_sha, + recipe=asdict(cfg), + )) + history = [] + previous_shard = None + preflight_scenarios = SP.expansion_scenarios(1, smoke=cfg.smoke) + preflight_replicas = [ + BX.Replica( + scenario_id, gamma, + n_ped=environment["n_ped"], + ped_speed_range=tuple(environment["ped_speed_range"]), + ) + for scenario_id in preflight_scenarios for gamma in SP.GAMMAS + ] + ell0, ell, ell_preflight = OE._initial_lengthscale( + policy, preflight_replicas, cfg, device, + ) + with ProcessPoolExecutor(max_workers=cfg.verifier_workers) as executor: + for round_i in range(1, cfg.rounds + 1): + round_start = time.perf_counter() + scenarios = SP.expansion_scenarios(round_i, smoke=cfg.smoke) + replicas = [ + BX.Replica( + scenario_id, gamma, + n_ped=environment["n_ped"], + ped_speed_range=tuple(environment["ped_speed_range"]), + ) + for scenario_id in scenarios for gamma in SP.GAMMAS + ] + if len(replicas) != 56: + raise RuntimeError("offline macro-round requires 56 episodes") + policy.eval() + phi_policy = copy.deepcopy(policy).eval() + for parameter in phi_policy.parameters(): + parameter.requires_grad_(False) + gp, gp_ids, gp_selection = OE.gp_from_previous( + phi_policy, previous_shard, round_i=round_i, ell=ell, + cap=OE.CAP, lam=cfg.gp_lam, phi_s=cfg.phi_s, device=device, + seed=cfg.seed + round_i * 101, + ) + beta, calibrated_ess = OE._calibrate_beta( + phi_policy, gp, replicas, cfg, device, round_i=round_i, + ) + shard = OS.ExecutedRoundShard(round_i) + gather = OE.gather_offline_round( + policy, phi_policy, gp, beta, replicas, cfg, shard, device, + executor, round_i=round_i, + ) + shard_path = os.path.join( + outdir, "round_shards", f"round_{round_i:02d}.pt", + ) + shard_manifest = shard.save(shard_path) + replay_start = time.perf_counter() + replay = AUG.replay_with_mode( + policy, optimizer, shard, + mode=cfg.replay_mode, alpha=cfg.alpha, + exposure_epochs=cfg.exposure_epochs, batch=cfg.batch, + device=device, seed=cfg.seed + round_i * 1_000_003, + executor=executor, + ) + gather["timers"]["replay"] = time.perf_counter() - replay_start + if BS.module_sha256(policy.enc_grid) != visual_encoder_sha: + raise RuntimeError("visual encoder SHA changed") + checkpoint_path = os.path.join(outdir, f"round_{round_i:02d}.pt") + BX._save_checkpoint(policy, checkpoint_path, dict( + round=round_i, experiment=cfg.arm_name, + source_checkpoint=checkpoint, source_sha256=checkpoint_sha, + encoder_sha256=visual_encoder_sha, recipe=asdict(cfg), + ell=ell, ell0=ell0, cap=OE.CAP, beta=float(beta), + )) + record = dict( + round=round_i, experiment=cfg.arm_name, + scenarios=list(map(int, scenarios)), + environment=environment, beta=float(beta), + calibrated_normalized_ess_over_remaining=float(calibrated_ess), + verifier=SM.verifier_manifest(), + gp_buffer_ids=gp_ids, gp_selection=gp_selection, + gp=gp.diagnostics(), gather=gather, replay=replay, + shard=shard_manifest, + checkpoint=os.path.abspath(checkpoint_path), + checkpoint_sha256=OS.sha256_file(checkpoint_path), + wall_seconds=time.perf_counter() - round_start, + ) + history.append(record) + with open(os.path.join(outdir, "metrics.jsonl"), "a") as stream: + stream.write(json.dumps(record, allow_nan=False) + "\n") + print(json.dumps(dict( + round=round_i, experiment=cfg.arm_name, + D=shard_manifest["D"], Dplus=shard_manifest["Dplus"], + Dminus=shard_manifest["Dminus"], beta=float(beta), + replay_mode=cfg.replay_mode, + Adam_steps=int(replay["optimizer_steps"]), + wall_seconds=record["wall_seconds"], + )), flush=True) + previous_shard = shard + + manifest = dict( + status="CLAUDE_SFM_B1_OFFLINE_EXT_COMPLETE", + experiment=cfg.arm_name, + scientific_role="offline_expansion_data_collector_not_safe_controller", + recipe=asdict(cfg), + replay_rules=AUG.declared_rules(), + constants=dict( + ell=ell, ell0=ell0, ell_preflight=ell_preflight, + gp_buffer_cap=OE.CAP, gp_lambda=OE.GP_LAMBDA, + expected_checkpoint_sha256=OE.EXPECTED_CHECKPOINT_SHA256, + replay_window_rounds=1, + ), + source=OE._source(), + source_checkpoint=checkpoint, + source_checkpoint_sha256=checkpoint_sha, + environment=environment, + frozen_parameters=frozen_parameters, + visual_encoder_sha=visual_encoder_sha, + history=history, + ) + OE._write_json(os.path.join(outdir, "COMPLETE.json"), manifest) + return manifest + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--checkpoint", required=True) + parser.add_argument("--outdir", required=True) + parser.add_argument("--alpha", type=float, required=True) + parser.add_argument("--exposure-epochs", type=int, required=True) + parser.add_argument( + "--selector", choices=OE.EXECUTION_SELECTORS, default="margin", + ) + parser.add_argument("--rounds", type=int, default=10) + parser.add_argument("--lr", type=float, default=1.0e-4) + parser.add_argument("--ess-target", type=float, default=0.5) + parser.add_argument( + "--replay-mode", choices=AUG.REPLAY_MODES, default="original", + ) + parser.add_argument("--verifier-workers", type=int, default=8) + parser.add_argument("--seed", type=int, default=20260724) + parser.add_argument("--device", default="cuda") + parser.add_argument("--smoke", action="store_true") + parser.add_argument("--tag", default="ext") + args = parser.parse_args(argv) + cfg = ExtConfig( + alpha=args.alpha, exposure_epochs=args.exposure_epochs, + selector=args.selector, rounds=args.rounds, lr=args.lr, + ess_target=args.ess_target, replay_mode=args.replay_mode, + verifier_workers=args.verifier_workers, seed=args.seed, + smoke=args.smoke, tag=args.tag, + ) + run(args.checkpoint, args.outdir, cfg, device=args.device) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/claude_paper_trends.py b/overnight_run_07_12_sfm/claude_paper_trends.py new file mode 100644 index 0000000..f78c188 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_paper_trends.py @@ -0,0 +1,105 @@ +"""Convert raw offline evaluator metrics into the paper trends contract. + +Reads one or more ``raw_m{M}_offline_metrics.json`` files produced by +``sfm_b1_offline_eval.py`` (each containing per-round records with the full +per-episode rows) and writes the per-(round,gamma) JSONL consumed by +``safe_flow_expansion@87063d3:scripts/paper_b1_margin50_trends.py``: + + {"round": r, "gamma": g, "m": M, "temp": 1.0, + "CR": {"mean": .., "se": ..}, "v_safe": {"mean": .., "se": ..}, + "clearance": {"mean": .., "se": ..}, "time": {"mean": .., "se": ..}} + +Standard errors are computed numerically from the stored rows (binomial SE +for CR; sample SE of per-trajectory validity fractions; sample SE over +successful episodes for clearance and time) and stored alongside the means, +as the study contract requires. The band semantics of the paper script are +unchanged: it applies Wilson intervals to CR/v_safe from (mean, m) and +mean +/- 1.96*se to clearance/time. +""" +from __future__ import annotations + +import argparse +import json +import math +import os +import subprocess +import sys + + +def _se_binomial(p, n): + return math.sqrt(max(p * (1.0 - p), 0.0) / n) if n else float("nan") + + +def _mean_se(values): + finite = [float(v) for v in values if v is not None] + if not finite: + return None, None + mean = sum(finite) / len(finite) + if len(finite) < 2: + return mean, 0.0 + var = sum((v - mean) ** 2 for v in finite) / (len(finite) - 1) + return mean, math.sqrt(var / len(finite)) + + +def convert(metrics_paths, jsonl_path): + rows_out = [] + for path in metrics_paths: + with open(path) as stream: + payload = json.load(stream) + for record in payload["records"]: + cell = record["cell"] + for gamma_key in cell["summary"]["per_gamma"]: + gamma = float(gamma_key) + rows = [ + row for row in cell["rows"] + if float(row["gamma"]) == gamma + ] + m = len(rows) + cr = sum(bool(r["collision"]) for r in rows) / m + v_mean, v_se = _mean_se([r["validity"] for r in rows]) + c_mean, c_se = _mean_se( + [r["successful_clearance"] for r in rows], + ) + t_mean, t_se = _mean_se([r["time_to_goal"] for r in rows]) + rows_out.append(dict( + round=int(record["round"]), gamma=gamma, m=m, temp=1.0, + CR=dict(mean=cr, se=_se_binomial(cr, m)), + v_safe=dict(mean=v_mean, se=v_se), + clearance=dict(mean=c_mean, se=c_se), + time=dict(mean=t_mean, se=t_se), + )) + rows_out.sort(key=lambda r: (r["round"], r["gamma"])) + os.makedirs(os.path.dirname(os.path.abspath(jsonl_path)), exist_ok=True) + with open(jsonl_path, "w") as stream: + for row in rows_out: + stream.write(json.dumps(row, allow_nan=False) + "\n") + return rows_out + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--metrics", nargs="+", required=True, + help="raw_m*_offline_metrics.json files (rounds merge)") + parser.add_argument("--label", required=True) + parser.add_argument("--outdir", required=True) + parser.add_argument("--stem", default="b1_margin50_metric_trends") + parser.add_argument( + "--paper-script", + default=("/home/dohyun/projects/safe_flow_expansion-claude-plot-" + "87063d3/scripts/paper_b1_margin50_trends.py"), + ) + args = parser.parse_args(argv) + outdir = os.path.abspath(args.outdir) + jsonl_path = os.path.join(outdir, f"{args.label}_trends_rows.jsonl") + rows = convert(args.metrics, jsonl_path) + print(f"{len(rows)} rows -> {jsonl_path}") + command = [ + sys.executable, args.paper_script, + "--arm", f"{args.label}={jsonl_path}", + "--outdir", outdir, "--stem", args.stem, + ] + subprocess.run(command, check=True) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/claude_priv_online.py b/overnight_run_07_12_sfm/claude_priv_online.py new file mode 100644 index 0000000..96374dd --- /dev/null +++ b/overnight_run_07_12_sfm/claude_priv_online.py @@ -0,0 +1,247 @@ +"""Phase 1+2: online privileged-controller qualification and faithful D_MPC. + +Phase 1 runs the EXACT historical privileged deterministic SFM controller +(``privileged_sfm_config`` @0e0eca2 driving the unmodified +``kazuki_sfm_deploy``: flow/guidance/MPPI nominal pool, brake/goal/avoidance/ +constant-acceleration templates, privileged candidate-specific SFM +simulation, H10 hard margin + recoverability, H20 viability/escalation, +stagnation/escape, always_select=True, execute-one-action-and-replan) on the +fixed CRN teacher bank, reporting SR/CR/timeout/executed-window Validity/ +clearance/time. Compact SOCP is audit-only. + +Phase 2 keeps ONLY successful controller episodes and builds ``D_MPC`` from +the sequence of ACTUALLY EXECUTED first actions: contexts are reconstructed +offline exactly as evaluation-time contexts (deterministic pedestrian replay +verified against the rollout's stored arrays, fail-closed), H10 windows are +sliding windows over executed actions, terminal prefixes shorter than H10 +are omitted, and every window carries step-filter/selection/clearance +provenance plus an audit-only compact-SOCP label. D_MPC never enters D, +D+/D-, the GP, acquisition, or beta calibration. +""" +from __future__ import annotations + +import argparse +from concurrent.futures import ProcessPoolExecutor +import json +import multiprocessing as mp +import os + +import numpy as np +import torch + +import _paths # noqa: F401 + + +def _rollout_gamma(payload): + (checkpoint, ep0, m_per_gamma, gamma, device) = payload + import grid_feats as GF + import grid_policy_sfm as GPS + import claude_mpc_pool as MP + import sfm_b1_offline_eval as OE_EVAL + import sfm_hp_history as HH + import sfm_kazuki as KZ + import sfm_metrics2 as SM + import sfm_protocol as SP + import sfm_scene as SS + + environment = SS.scene_profile("double_density_velocity_ood") + policy, _ = GPS.load_sfm_policy(checkpoint, device=device) + policy.eval() + config = MP.privileged_sfm_config() + rows, lineages = [], [] + for episode in range(int(ep0), int(ep0) + int(m_per_gamma)): + rollout = KZ.kazuki_sfm_deploy( + policy, episode, float(gamma), cfg=config, + n_ped=environment["n_ped"], T=int(SP.T), reach=0.5, + device=device, + ped_speed_range=tuple(environment["ped_speed_range"]), + sample_seed=700_000, collect_diagnostics=True, + ) + success = bool(rollout["success"]) + steps = int(rollout["steps"]) + row = { + "episode": int(episode), "gamma": float(gamma), + "status": ("success" if success else + "collision" if rollout["collision"] else "timeout"), + "success": success, + "collision": bool(rollout["collision"]), + "timeout": bool(not success and not rollout["collision"]), + "steps": steps, + "time_to_goal": steps * SS.DT if success else None, + "min_clearance": float(rollout["min_clear"]), + "successful_clearance": ( + float(rollout["min_clear"]) if success else None + ), + "states": np.asarray(rollout["states"], np.float32), + "controls": np.asarray(rollout["controls"], np.float32), + "ped_xy": np.asarray(rollout["peds"], np.float32), + "ped_vel": np.asarray(rollout["ped_vels"], np.float32), + } + validity = OE_EVAL._verify_executed_episode(row) + compact = {k: row[k] for k in row if k not in ( + "states", "controls", "ped_xy", "ped_vel", + )} + compact.update(validity) + rows.append(compact) + if not success or steps < 10: + continue + # --- Phase 2: faithful context reconstruction (fail-closed) --- + humans = SS.make_humans( + episode, 0, environment["n_ped"], + tuple(environment["ped_speed_range"]), + ) + state = np.zeros(4, np.float32) + history = HH.HpHistory() + contexts = [] + trace = list(rollout.get("trace") or ()) + for t in range(steps): + ped_xy, ped_vel = SS.collect_humans(humans) + if not ( + np.allclose(ped_xy, row["ped_xy"][t], atol=1e-4) + and np.allclose(state, row["states"][t], atol=1e-4) + ): + raise RuntimeError( + f"provenance mismatch reconstructing s{episode} " + f"g{gamma} t{t}" + ) + obstacles = np.concatenate([ + ped_xy, + np.full((len(ped_xy), 1), SS.R_PED, np.float32), + ], axis=1) + hp10 = history.append(torch.as_tensor(GF.axis_grid( + state[:2], obstacles, 0.0, R=SS.R_SENSE, + sensing=SS.R_SENSE, + ))) + low = GF.low5(state, SS.GOAL, float(gamma)) + hist = GF.hist_pad( + row["controls"][max(0, t - 16):t] + if t else np.zeros((0, 2)), 16, + ) + diag = trace[t] if t < len(trace) else {} + step_filter = diag.get("output_filter") or {} + contexts.append(dict( + step=t, state=state.copy(), + hp10=hp10.numpy().astype(np.float32), + low5=np.asarray(low, np.float32), + hist=np.asarray(hist, np.float32), + ped_xy=ped_xy.copy(), ped_vel=ped_vel.copy(), + selection_reason=step_filter.get("selection_reason"), + filter_feasible=step_filter.get("filter_feasible"), + horizon_clear=step_filter.get("selected_horizon_clear"), + )) + action = row["controls"][t] + state = state.copy() + state[:2] += SS.DT * state[2:4] + 0.5 * SS.DT ** 2 * action + state[2:4] += SS.DT * action + SS.advance_humans(humans, state) + windows = [] + for start in range(steps - 10 + 1): + controls = row["controls"][start:start + 10] + audit = SM.verify_query( + contexts[start]["state"], controls, + contexts[start]["ped_xy"], contexts[start]["ped_vel"], + float(gamma), + ) + windows.append(dict( + start=start, + controls=np.asarray(controls, np.float32), + context=contexts[start], + socp_audit_only=dict( + resolved=bool(audit.get("resolved")), + y=int(audit.get("y", 0)) if audit.get("resolved") else None, + ), + )) + lineages.append(dict( + episode=int(episode), gamma=float(gamma), steps=steps, + windows=windows, + )) + return rows, lineages + + +def run(args): + outdir = os.path.abspath(args.outdir) + if os.path.exists(outdir): + raise FileExistsError(outdir) + os.makedirs(outdir) + import sfm_b1_offline_eval as OE_EVAL + import sfm_b1_offline_store as OS + import sfm_protocol as SP + + payloads = [ + (os.path.abspath(args.checkpoint), int(args.ep0), + int(args.m_per_gamma), float(gamma), args.device) + for gamma in SP.GAMMAS + ] + context = mp.get_context("spawn") + with ProcessPoolExecutor( + max_workers=min(7, int(args.workers)), mp_context=context, + ) as executor: + results = list(executor.map(_rollout_gamma, payloads)) + rows = [row for cell_rows, _ in results for row in cell_rows] + lineages = [l for _, cell_lineages in results for l in cell_lineages] + summary = OE_EVAL.summarize(rows, seed=int(args.ep0)) + OE_EVAL._assert_zero_verifier_errors(summary) + windows_total = sum(len(l["windows"]) for l in lineages) + per_gamma = {} + audit_positive = 0 + for lineage in lineages: + key = str(lineage["gamma"]) + per_gamma[key] = per_gamma.get(key, 0) + len(lineage["windows"]) + audit_positive += sum( + 1 for w in lineage["windows"] + if w["socp_audit_only"]["y"] == 1 + ) + buffer = dict( + status="CORRECTED_D_MPC_BUFFER_COMPLETE", + checkpoint=os.path.abspath(args.checkpoint), + checkpoint_sha256=OS.sha256_file(args.checkpoint), + bank=dict(ep0=int(args.ep0), m_per_gamma=int(args.m_per_gamma)), + controller="privileged_sfm_config @0e0eca2 via unmodified kazuki_sfm_deploy", + retention="successful episodes only; H10 windows over EXECUTED actions; terminal prefixes < H10 omitted", + separation="never enters D, D+/D-, GP, acquisition, or beta calibration", + lineages=lineages, + n_lineages=len(lineages), + n_windows=windows_total, + windows_per_gamma=per_gamma, + socp_audit_positive_fraction=( + audit_positive / windows_total if windows_total else None + ), + ) + torch.save(buffer, os.path.join(outdir, "D_MPC_corrected.pt")) + report = dict( + status="PRIV_ONLINE_QUALIFICATION_COMPLETE", + controller_summary=summary, + rows=rows, + n_lineages=len(lineages), + n_windows=windows_total, + windows_per_gamma=per_gamma, + socp_audit_positive_fraction=buffer[ + "socp_audit_positive_fraction" + ], + ) + OE_EVAL._write_json(os.path.join(outdir, "qualification.json"), report) + pooled = summary["pooled"] + print(json.dumps(dict( + SR=pooled["SR"], CR=pooled["CR"], timeout=pooled["timeout"], + Validity=pooled["Validity"]["mean"], + clearance=pooled["successful_clearance"]["mean"], + time=pooled["successful_time_to_goal"]["mean"], + lineages=len(lineages), windows=windows_total, + audit_positive_fraction=buffer["socp_audit_positive_fraction"], + ), allow_nan=False), flush=True) + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--checkpoint", required=True) + parser.add_argument("--ep0", type=int, required=True) + parser.add_argument("--m-per-gamma", type=int, default=10) + parser.add_argument("--workers", type=int, default=7) + parser.add_argument("--device", default="cuda:0") + parser.add_argument("--outdir", required=True) + args = parser.parse_args(argv) + run(args) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/claude_r1_anchor.py b/overnight_run_07_12_sfm/claude_r1_anchor.py new file mode 100644 index 0000000..6894d16 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_r1_anchor.py @@ -0,0 +1,202 @@ +"""Freeze the r1 self-anchor buffer for the stable-continuation study. + +Runs the immutable r1 checkpoint RAW (temperature 1, NFE 8, CRN latents) on +the private anchor bank, keeps only SUCCESSFUL episodes, exact-verifies every +executed sliding window of full length H_t=10, and stores the y=1 windows +with their exact gathering-format contexts (hp10/low5/hist/state/ped_xy/ +ped_vel). This is self-replay of the policy's own certified successful +behavior — not expert data. The buffer is frozen once, before round 1, with +a deterministic gamma-balanced cap (order by (episode, step), first 200 per +gamma). +""" +from __future__ import annotations + +import argparse +from concurrent.futures import ProcessPoolExecutor +import json +import multiprocessing as mp +import os + +import numpy as np +import torch + +import _paths # noqa: F401 +import grid_feats as GF +import grid_policy_sfm as GPS +import sfm_b1_eval as BE +import sfm_b1_offline_store as OS +import sfm_hp_history as HH +import sfm_metrics2 as SM +import sfm_protocol as SP +import sfm_scene as SS + +CAP_PER_GAMMA = 200 +H = 10 + + +@torch.no_grad() +def _rollout_with_contexts(policy, episode, gamma, noise, device, environment): + humans = SS.make_humans( + int(episode), 0, environment["n_ped"], + tuple(environment["ped_speed_range"]), + ) + state = np.zeros(4, np.float32) + history = HH.HpHistory() + controls_list, contexts = [], [] + status = None + minimum_clearance = float("inf") + for step in range(int(SP.T)): + ped_xy, ped_vel = SS.collect_humans(humans) + clearance = float( + np.linalg.norm(ped_xy - state[:2][None], axis=1).min() - SS.R_PED + ) if len(ped_xy) else float("inf") + minimum_clearance = min(minimum_clearance, clearance) + if clearance < 0.0: + status = "collision" + break + if float(np.linalg.norm(state[:2] - SS.GOAL)) < 0.5: + status = "success" + break + obstacles = np.concatenate([ + ped_xy, np.full((len(ped_xy), 1), SS.R_PED, np.float32), + ], axis=1) + raw_grid = torch.as_tensor(GF.axis_grid( + state[:2], obstacles, 0.0, R=SS.R_SENSE, sensing=SS.R_SENSE, + )) + hp10 = history.append(raw_grid) + low = torch.as_tensor(GF.low5(state, SS.GOAL, gamma)) + hist = torch.as_tensor(GF.hist_pad( + np.asarray(controls_list[-16:]) + if controls_list else np.zeros((0, 2)), 16, + )) + ctx = policy.ctx_from( + hp10[None].float().to(device), low[None].float().to(device), + hist[None].float().to(device), + ) + window = BE.integrate_latents( + policy, + torch.as_tensor(noise[step][None], device=device), + ctx, nfe=8, + ).reshape(H, 2).cpu().numpy().astype(np.float32) + contexts.append(dict( + step=int(step), state=state.copy(), + hp10=hp10.numpy().astype(np.float32), + low5=low.numpy().astype(np.float32), + hist=hist.numpy().astype(np.float32), + ped_xy=ped_xy.copy(), ped_vel=ped_vel.copy(), + )) + action = window[0] + controls_list.append(action) + state[:2] = state[:2] + SS.DT * state[2:4] + 0.5 * SS.DT ** 2 * action + state[2:4] = state[2:4] + SS.DT * action + SS.advance_humans(humans, state) + if status is None: + status = "timeout" + return dict( + episode=int(episode), gamma=float(gamma), status=status, + steps=len(controls_list), + controls=np.asarray(controls_list, np.float32), + contexts=contexts, min_clearance=float(minimum_clearance), + ) + + +def collect(checkpoint, outpath, *, ep0, m_per_gamma, noise_seed, device, + workers): + policy, _ = GPS.load_sfm_policy(checkpoint, device=device) + policy.eval() + environment = SS.scene_profile("double_density_velocity_ood") + generator = np.random.default_rng(int(noise_seed)) + noise = generator.standard_normal( + (len(SP.GAMMAS), int(m_per_gamma), int(SP.T), int(policy.d)), + dtype=np.float32, + ) + rollouts = [] + for gamma_index, gamma in enumerate(SP.GAMMAS): + for rollout_index in range(int(m_per_gamma)): + rollouts.append(_rollout_with_contexts( + policy, int(ep0) + rollout_index, float(gamma), + noise[gamma_index, rollout_index], device, environment, + )) + successes = [r for r in rollouts if r["status"] == "success"] + tasks, meta = [], [] + for run_index, run in enumerate(successes): + n = run["steps"] + for start in range(0, n - H + 1): + tasks.append(( + len(meta), 0, run["contexts"][start]["state"], + run["controls"][start:start + H], + run["contexts"][start]["ped_xy"], + run["contexts"][start]["ped_vel"], run["gamma"], + )) + meta.append((run_index, start)) + context_mp = mp.get_context("spawn") + with ProcessPoolExecutor( + max_workers=int(workers), mp_context=context_mp, + ) as executor: + results = {i: r for i, _, r in executor.map(SM.verify_in_worker, tasks)} + per_gamma = {} + for task_index, (run_index, start) in enumerate(meta): + result = results[task_index] + if not result.get("resolved") or int(result.get("y", 0)) != 1: + continue + run = successes[run_index] + per_gamma.setdefault(round(run["gamma"], 8), []).append(dict( + episode=run["episode"], gamma=run["gamma"], step=start, + context=run["contexts"][start], + controls=run["controls"][start:start + H].copy(), + )) + records = [] + audit_counts = {} + for gamma in sorted(per_gamma): + rows = sorted(per_gamma[gamma], key=lambda r: (r["episode"], r["step"])) + kept = rows[:CAP_PER_GAMMA] + audit_counts[str(gamma)] = dict( + certified_available=len(rows), kept=len(kept), + ) + records.extend(kept) + payload = dict( + status="R1_ANCHOR_BUFFER_FROZEN", + checkpoint=os.path.abspath(checkpoint), + checkpoint_sha256=OS.sha256_file(checkpoint), + bank=dict(ep0=int(ep0), m_per_gamma=int(m_per_gamma), + noise_seed=int(noise_seed)), + cap_per_gamma=CAP_PER_GAMMA, + rollout_outcomes={ + status: sum(r["status"] == status for r in rollouts) + for status in ("success", "collision", "timeout") + }, + windows_verified=len(tasks), + per_gamma=audit_counts, + n_records=len(records), + records=records, + ) + torch.save(payload, outpath) + summary = {k: payload[k] for k in ( + "status", "checkpoint_sha256", "rollout_outcomes", + "windows_verified", "per_gamma", "n_records", + )} + with open(outpath + ".summary.json", "w") as stream: + json.dump(summary, stream, indent=1) + print(json.dumps(summary, indent=1)) + return payload + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--checkpoint", required=True) + parser.add_argument("--out", required=True) + parser.add_argument("--ep0", type=int, default=380_000) + parser.add_argument("--m-per-gamma", type=int, default=12) + parser.add_argument("--noise-seed", type=int, default=20_260_739) + parser.add_argument("--device", default="cuda:0") + parser.add_argument("--workers", type=int, default=16) + args = parser.parse_args(argv) + collect( + args.checkpoint, args.out, ep0=args.ep0, + m_per_gamma=args.m_per_gamma, noise_seed=args.noise_seed, + device=args.device, workers=args.workers, + ) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/claude_recipe_study/AFE_FIDELITY_DELIVERY.json b/overnight_run_07_12_sfm/claude_recipe_study/AFE_FIDELITY_DELIVERY.json new file mode 100644 index 0000000..2606e94 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/AFE_FIDELITY_DELIVERY.json @@ -0,0 +1,7 @@ +{"status":"AFE_FIDELITY_DELIVERY_COMPLETE","frozen_sha":"0c52010", +"phase0":"P(pool has SOCP positive)=.938, mean multiplicity 26.9, yet full-pool certified execution NVP .986 (CR 0.000) and B8 NVP 1.0: closed-loop verifier-support intersection is the binding limitation (pre-registered interpretation), not CFM, controller quality, or acquisition ranking", +"screening":"8-arm Pareto table in SCREEN.json; promoted ess0p5_U_d1 + ess0p25_U_d1 (U objective, lr 1e-5, 4 whole-dataset steps)", +"expansion":"both recipes 10 certified macro-rounds from immutable r1; best per recipe = round 1 (EXPANSION_M10.json)", +"m50":"r1 .303 CR / .730 V; both recipes .289 CR / .737 V; winner ess0p25_U_d1 r1", +"final_m100_untouched_480000":{"r1":{"SR":0.6857,"CR":0.2771,"V":0.7278,"clr":0.1223,"t":11.006},"winner":{"SR":0.6914,"CR":0.2714,"V":0.7316,"clr":0.1238,"t":11.134}}, +"verdict":"small consistent improvement on every axis (dCR -.006, dV +.004, dSR +.006, dclr +.002) but WITHIN NOISE at n=700; no 97-99% transfer claim; the certified-set CFM transfers a real but epsilon-scale fraction of controller support into the raw flow"} diff --git a/overnight_run_07_12_sfm/claude_recipe_study/AFE_FIDELITY_PREREGISTRATION.json b/overnight_run_07_12_sfm/claude_recipe_study/AFE_FIDELITY_PREREGISTRATION.json new file mode 100644 index 0000000..946115b --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/AFE_FIDELITY_PREREGISTRATION.json @@ -0,0 +1,25 @@ +{ + "status": "PREDECLARED_BEFORE_ANY_RUN", + "declared_at": "2026-07-29T01:10:00-07:00", + "branch": "agent/claude-sfm-afe-fidelity-controller-benchmark-20260729 from frozen 736fa6c; artifacts under /data3/research1/claude_sfm_afe_fidelity_controller_/", + "golden_rule": "exact full-H moving-pedestrian SOCP is the ONLY execution and replay authority: enter D+/execute/train iff resolved AND y=1 AND full_h AND terminal_step=10; privileged SFM simulation generates and ranks but never certifies; always_select disabled for certified benchmarks; NVP terminates fail-closed; no relabeling", + "controller_score_J": "faithful replication of the committed exact_sfm_horizon_filter_action score with per-gamma weights: J = gsw*goal_cost + .015*mean(u^2) + .02*||u0-u0_nominal||^2 - cw*min(clear,1) + clearance_target_cost, goal_cost = .04*arrival if reached<=H10 else ||terminal-GOAL||; computed from the privileged simulation diagnostics; used for F/B8 execution ranking and as J_ctrl in objective G", + "banks": { + "phase0": {"ep0": 460000, "m_per_gamma": 10, "role": "O/F/B8 support benchmark, fixed CRN"}, + "gathering": "standard expansion bank; macro-round k gathers expansion_scenarios(1+k) (56 lineages)", + "m10_eval": {"ep0": 390000, "noise_seed": 20260740, "m_per_gamma": 10, "role": "identical CRN raw screening bank across arms and rounds (development-only)"}, + "m50": {"ep0": 470000, "noise_seed": 20260750, "m_per_gamma": 50}, + "m100": {"ep0": 480000, "noise_seed": 20260751, "m_per_gamma": 100, "role": "untouched until the single winner is frozen"} + }, + "afe": "one keyed x0~N(0,I) per proposal incl. deterministic templates (persisted); z=normalize(phi_s_from_x0((1-s)x0+sU,s,c)), s=.9; RBF GP rebuilt once per macro-round from previous round's exact-SOCP-positive queries (cap 512, gamma-equal quota, ell=0.3259987743470518 lambda=.01 from the authenticated preflight, fail-closed recalibration only on numerical condition), adaptive beta for ESS in {.25,.50}, B=8 without replacement, no within-round GP update; round-1 previous = the archived B9 round-1 shard (certified executed positives with stored x0)", + "store": "certified-query shard: D = every resolved B query, D+ = ALL exact positives per context (set-valued); episodes terminate at NVP", + "objectives": { + "U": "L_U = sum_c mass(c) * mean over P_c of L_CFM (== hierarchy mass gamma->cell->context->certified query)", + "G": "q(U|c) ∝ exp(-(J-Jmin)/tau_J), tau_J calibrated once per round so within-context median ESS/|P_c| = 0.5 (bisection, no manual temperature)", + "replay": "corrected whole-dataset accumulation, ONE Adam step per epoch, W=2 fail-closed (round 1 window = {B9 r1 shard, round-1 certified shard}; k>=2 = {k-1,k}), production cfm_loss weighting (chunk weights b*w_i summed), frozen encoder, positives only; no teacher/demo/fallback/prox/anchor/rollback/recovery/uncertified data" + }, + "screening": "2 gathers from r1 (ESS .25/.50, shared across objectives/doses) x {U,G} x {d1: lr1e-5 4 steps, d2: lr3e-5 4 steps} = 8 one-round arms on the M10 bank; promote every Pareto-nondominated arm then at most best two by (CR, -Validity, -clearance, -SR, timeout, time); full Pareto table always reported", + "expansion": "two promoted recipes x 10 macro-rounds from immutable r1; checkpoints every round; M10 at rounds 0,1,2,5,10; per-recipe best by M10 only; both to disjoint M50; single winner to untouched M100", + "compute": "GPU3 exclusively verified now; GPU1 re-checked by UUID before the expansion stage and used only if exclusively free; shared read-only verifier cache keyed by (context-hash, U-hash, gamma, verifier version)", + "interpretation": "four methods kept distinct (unverified oracle / full-pool SOCP-gated / budgeted B8 AFE / raw expanded flow); no 97-99% transfer claim unless the raw M100 achieves it; poor-liveness full pool => verifier-support gap binding; full works+B8 fails => acquisition binding; B8 works+raw fails => CFM transfer binding" +} diff --git a/overnight_run_07_12_sfm/claude_recipe_study/CORRECTED_DISTILL_PREREGISTRATION.json b/overnight_run_07_12_sfm/claude_recipe_study/CORRECTED_DISTILL_PREREGISTRATION.json new file mode 100644 index 0000000..f3a878a --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/CORRECTED_DISTILL_PREREGISTRATION.json @@ -0,0 +1,40 @@ +{ + "status": "PREDECLARED_BEFORE_ANY_RUN", + "declared_at": "2026-07-28T09:10:00-07:00", + "branch": "agent/claude-sfm-corrected-privileged-distill-20260728 (code pushed before running; master/Codex branches/PR3/prior checkpoints and artifacts untouched)", + "artifact_root": "/data3/research1/claude_sfm_corrected_privileged_distill_/ where FROZEN_SHA = the pushed source commit that launches the runs", + "disclosed_prior_bug": "the last two studies' admissibility gates read records[0] (= r0) instead of the r1 record from the dev-baseline metrics file; re-checking the archived tables against the CORRECT r1 gates still yields zero admissible candidates in both studies (timeout/SR/clearance still fail), so no conclusion flips, but the bar actually applied was r0's (CR<=.300/V>=.638/clr>=.119 with timeout<=.050) rather than r1's (CR<=.186/V>=.752/clr>=.144 with timeout<=.093); fixed here by a label-based reader", + "phase1_online_qualification": { + "controller": "EXACT historical privileged deterministic SFM controller: sfm_adhoc_controller_compare.privileged_sfm_config() (verbatim import @0e0eca2) driving sfm_kazuki.kazuki_sfm_deploy with exact_sfm_horizon_filter_action and its full candidate family (flow/guidance/MPPI nominal pool + brake + 12 goal + 18 avoidance + 25 constant-acceleration), privileged candidate-specific SFM simulation, H10 hard-margin/recoverability, H20 viability/escalation, stagnation/escape, always_select=True, execute-one-action-and-replan; compact SOCP audit-only, never execution authority", + "prior": "immutable r1 checkpoint (the policy the study intends to improve); the historical 98% figure belonged to a different scene and prior and is NOT claimed to transfer", + "bank": {"ep0": 440000, "m_per_gamma": 10, "role": "fixed CRN M10/gamma qualification AND phase-2 teacher generation (same fixed scenarios); disjoint from every evaluation bank"}, + "report": "raw r1 and online controller separately: SR, CR, timeout, executed-window Validity, successful clearance, successful time-to-goal" + }, + "phase2_teacher_dataset": { + "retention": "successful controller episodes only; every context + the actually executed first action; H10 windows reconstructed post-episode from consecutive EXECUTED actions only (never the nine unexecuted actions of a proposed plan); terminal prefixes < H10 omitted", + "contexts": "reconstructed offline exactly as evaluation-time contexts (HpHistory over axis_grid of deterministically replayed pedestrians; hist from executed controls; low5 from states+gamma) and verified against the rollout's stored pedestrian arrays (fail-closed on mismatch)", + "provenance_per_window": ["privileged-state/step-filter diagnostics of the executed step", "candidate family / selection_reason", "horizon clearance", "compact-SOCP audit label (audit-only)"], + "separation": "D_MPC never enters D, D+/D-, GP, acquisition, or beta calibration", + "render": "episode 250003, gamma=0.1, around t=55: all controller branches, feasible/rejected, selected, executed trajectory, raw r1 branches, post-distillation raw branches (PNG + MP4); explanatory evidence only" + }, + "phase3_fixes": { + "ordinary": "original B1 semantics via BS.signed_update on a fail-closed W=2 holder over BOTH declared recent shards (archived B9 round-1 shard + the newly gathered round-2 shard); one full-dataset accumulated gradient and exactly one Adam step per epoch; e10 would mean 10 Adam steps; abort rather than W=1", + "teacher": "L_T = sum_i m_i L_CFM(i) with normalized hierarchy masses; per chunk of size b pass weights=b*m_i to the production FlowPolicy.cfm_loss and SUM returned chunk losses (never multiply by chunk mass again); one accumulated Adam step per epoch; per-chunk deterministic seeding", + "test": "production GridSFMFlowPolicy test: chunked accumulated loss/gradient must numerically match an independent direct implementation that replays the identical per-chunk CFM noise (torch.randn then torch.rand, CPU) and computes sum_i m_i L_i analytically; gradient max-abs-diff tolerance 1e-5", + "baseline_reader": "label-based record selection; r1 reads the r1 record" + }, + "phase4_bounded_causal_study": { + "shared": "identical seeds, latents, ordering; evaluation bank = fixed raw temperature-1 M10/gamma ep0 390000 seed 20260740 (no guidance, no verifier acquisition, no controller templates, no temperature tuning)", + "arms": { + "A": "immutable r1, no update", + "B": "corrected ordinary B1 only: alpha=.01, ONE replay epoch (one accumulated Adam step) on W=2 {B9 r1 shard, new round-2 shard gathered from r1 with unchanged protocol}", + "C": "corrected teacher only, applied directly to r1 (before any ordinary continuation): lr in {1e-5, 3e-5} x epochs in {1, 4} = 4 arms", + "D": "best teacher dose (chosen from C on the M10 screening by the lexicographic safety rule among liveness-preserving rows) THEN one corrected ordinary epoch (B's exact settings)" + }, + "no_hard_gate_termination": "the complete Pareto table over (CR down, Validity up, clearance up, time down, SR, timeout) is reported for all arms", + "promotion_rule": "non-dominated candidates on (CR, -Validity, -clearance) that strictly improve at least one of CR/Validity/clearance versus r1 AND avoid catastrophic liveness loss, declared as SR >= SR_r1 - 0.15 and timeout <= timeout_r1 + 0.15 on the same bank, are promoted to the disjoint M50", + "m50": {"ep0": 450000, "noise_seed": 20260748, "m_per_gamma": 50, "role": "confirmation of promoted candidates vs r1 and r0"} + }, + "diagnostics_required": ["exact Adam-step count per arm", "parameter drift by module", "fixed-context first-action and H10 output RMSE (fixed probe contexts + fixed latents)", "teacher/ordinary-positive/D- gradient cosines", "controller-success lineage count", "teacher sample count and gamma balance", "compact-SOCP-positive fraction (audit only)", "raw control spread and route entropy"], + "scientific_rules": ["privileged teacher never called certified", "no transfer claim for the historical 98%", "no deterministic-distillation-failure claim unless the corrected teacher produces measurable parameter/output change", "stop and report on weighting mismatch, wrong Adam-step count, missing W=2 shard, provenance mismatch, or evaluation leakage"] +} diff --git a/overnight_run_07_12_sfm/claude_recipe_study/DELIVERY_COMPLETE.json b/overnight_run_07_12_sfm/claude_recipe_study/DELIVERY_COMPLETE.json new file mode 100644 index 0000000..84349f6 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/DELIVERY_COMPLETE.json @@ -0,0 +1,54 @@ +{ + "status": "CLAUDE_SFM_BEST_RECIPE_DELIVERY_COMPLETE", + "completed_at": "2026-07-26T18:20:00-07:00", + "headline": "A fixed Safe Flow Expansion recipe genuinely beat r0 on the disjoint M100 confirmation: CR -.100 [-.157,-.043], window Validity +.133 [+.114,+.153], successful clearance +.014 [+.005,+.023], SR +.080, at time +2.08 s [1.84,2.32] (paired scenario-cluster 95% CIs, 700 CRN rollouts/method, temperature 1.0, no tilt/fallback).", + "selected_recipe": { + "name": "margin_origrec_e100 (iteration 2, R_A)", + "execution_selector": "margin (max one-step nominal H_P margin among exact-certified B queries)", + "replay": "orig_plus_recovery: full executed D+ and D- plus exact-certified deterministic recovery positives (family v1) appended at hard contexts", + "alpha": 0.01, "exposure_epochs": 100, "lr": 1e-4, "ess_target": 0.5, + "rounds": 4, "seed": 20260724, + "selected_round": 1, + "selected_checkpoint": "/data3/research1/claude_sfm_best_recipe_f06e8dd/stageB/B9_margin_origrec_e100/round_01.pt", + "selected_checkpoint_sha256_prefix": "141e4ae6592bf73f" + }, + "provenance_chain": [ + "iteration 1 (frozen by predeclared Stage-B rule): cost/hard — honest NULL on its own M50 bank (STAGE_D_SELECTION.json); fully reported, nothing hidden", + "iteration 2 pre-registered BEFORE reading its selection bank (ITERATION2_PREREGISTRATION.json); recipe R_A chosen over v2 family by the declared U10-vs-U12 criterion", + "M50 selection on fresh bank 340000: r1 the only eligible round (ITERATION2_SELECTION.json)", + "M100 confirmation on untouched bank 330000: r0 / selected / locked Kazuki (stageE/)" + ], + "m100_confirmation_pooled": { + "r0": {"SR": 0.643, "CR": 0.353, "timeout": 0.004, "Validity": 0.603, "clearance": 0.108, "time": 8.76}, + "selected": {"SR": 0.723, "CR": 0.253, "timeout": 0.024, "Validity": 0.736, "clearance": 0.122, "time": 10.84}, + "kazuki": {"SR": 0.827, "CR": 0.173, "timeout": 0.000, "Validity": 0.353, "clearance": 0.168, "time": 4.17} + }, + "kazuki_note": "locked comparator (safe .3 / goal .5, Hp10 prior, native MPPI, no retuning): lower CR and faster via guidance+refinement, but its executed windows satisfy the exact GREEN certificate only 35% of the time versus 74% for our raw selected policy; it is a guided controller, not a raw flow", + "gamma_trend": "gamma=0.1 keeps the largest successful clearance (.148) and the longest time (13.4 s) for the selected checkpoint; mid-gamma flat within noise (per-gamma table stageE/m100_per_gamma_table.csv)", + "stability": "the recipe's gain is a round-1 phenomenon: round 2 fails the liveness gate (timeout .163 on the selection bank) and rounds 3-4 collapse to timeout; reported plainly, no monotonic-learning claim", + "local_repairs": "certified-recovery data demonstrably repairs the targeted failure modes: NVP contexts -25% and B-positive fraction .71->.84 on identical keyed latents; both diagnosed collision lineages became closed-loop successes under the selected checkpoint (mechanism/ figures)", + "integrity": [ + "temperature 1.0 everywhere, no per-gamma or global temperature tuning", + "no deterministic controller executed at evaluation; recovery data is replay-only", + "Kazuki untouched after declaration; checkpoint never selected from M100", + "iteration-1 null result and all failed arms reported (Stage-B table, Stage-D curve)", + "every synthetic positive carries an exact-verifier certificate audit row (iteration2/recovery_certificate_audit_B9.json)" + ], + "artifacts": { + "four_metric_plot_selected_recipe": "iteration2/paper_trends/b1_margin50_metric_trends.{png,pdf} (+ numeric SEs in margin_origrec_e100_trends_rows.jsonl)", + "four_metric_plot_iteration1": "stageD/paper_trends/b1_margin50_metric_trends.{png,pdf}", + "m50_selection": "iteration2/m50/", "m100_confirmation": "stageE/", + "paired_cis": "stageE/CONFIRMATION_ANALYSIS.json", + "per_gamma_tables": "stageE/m100_per_gamma_table.{csv,json}", + "population_counts": "iteration2/population_counts_B9.json", + "certificate_audit": "iteration2/recovery_certificate_audit_B9.json", + "mechanism_figures": "mechanism/ (+ artifact https://claude.ai/code/artifact/3c16fe2a-1450-4a8c-addf-466beaea1111)", + "baselines": "baselines/", + "manifest": "SHA256_MANIFEST.json (33 rehashed artifacts)", + "all_checkpoints": "stageB/B9_margin_origrec_e100/ (r0-r4), stageD/final_cost_hard/ (r0-r10)" + }, + "branches": { + "recipe_study": "agent/claude-sfm-best-recipe-20260726 (pushed)", + "mpc_distillation_followup": "agent/claude-sfm-mpc-distill-20260726 (superset; pushed; study running)" + } +} diff --git a/overnight_run_07_12_sfm/claude_recipe_study/EPISODE_BANKS.json b/overnight_run_07_12_sfm/claude_recipe_study/EPISODE_BANKS.json new file mode 100644 index 0000000..5151d4c --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/EPISODE_BANKS.json @@ -0,0 +1,44 @@ +{ + "status": "PREDECLARED_BEFORE_ANY_NEW_RUN", + "declared_at": "2026-07-26T02:33:00-07:00", + "declared_by": "claude agent/claude-sfm-best-recipe-20260726 @ f06e8ddc11fc7a1f5bada2cb2a587bff3ba4e424", + "scene_profile": "double_density_velocity_ood", + "environment": {"n_ped": 40, "ped_speed_range": [1.0, 2.0]}, + "gammas": [0.1, 0.2, 0.3, 0.4, 0.5, 0.7, 1.0], + "raw_evaluation": {"temperature": 1.0, "NFE": 8, "T": 180, "H": 10}, + "historical_banks_not_reused_for_selection": { + "expansion_training": {"ep0": 20000, "note": "8 scenarios/round; rounds 1-10 use 20000-20079; PRESERVED as the training expansion bank for all new training runs"}, + "codex_m10_screen": {"ep0": 260000, "noise_seed": 20260723, "role": "historical"}, + "codex_m50_selector_confirm": {"ep0": 270000, "noise_seed": 20260724, "role": "historical"}, + "codex_m100_final": {"ep0": 280000, "noise_seed": 20260725, "role": "historical; also reused here ONLY as the pre-registered baseline-reproduction bank (section 1 of the task), never for tuning or selection of the new recipe"} + }, + "new_private_banks_mutually_disjoint": { + "local_diagnosis": { + "ep0": 300000, "m_per_gamma": 8, "noise_seed": 20260726, + "role": "Stage A hard-episode mining and before/after single-update local checks (gathering-side); plus round-1 expansion episodes 20000-20007 which are training data" + }, + "anchor_raw": { + "ep0": 305000, "m_per_gamma": 12, "noise_seed": 20260727, + "role": "Stage A anchor bank: small fixed raw bank to detect global regressions caused by local repairs; never used for final selection" + }, + "qualification_raw": { + "ep0": 310000, "m_per_gamma": 25, "noise_seed": 20260728, + "role": "Stage B short-qualification fixed raw screening bank; candidate recipes are compared ONLY on this bank's raw policy metrics" + }, + "m50_checkpoint_selection": { + "ep0": 320000, "m_per_gamma": 50, "noise_seed": 20260729, + "role": "Stage D per-round fixed raw M50 CRN bank for the frozen final recipe; checkpoint selection uses ONLY this bank via the predeclared SELECTION_RULE.json" + }, + "m100_final_confirmation": { + "ep0": 330000, "m_per_gamma": 100, "noise_seed": 20260730, + "role": "Stage E disjoint confirmation of r0 vs selected expanded checkpoint vs locked Kazuki; never read before the checkpoint is frozen" + }, + "iteration2_m50_selection": { + "ep0": 340000, "m_per_gamma": 50, "noise_seed": 20260732, + "declared_at": "2026-07-26T13:15:00-07:00", + "role": "iteration-2 M50 checkpoint selection (see ITERATION2_PREREGISTRATION.json); declared because the iteration-1 selection bank (320000) was read for two declared diagnostic companion cells (B9 r1/r2) and is therefore retired for selection purposes" + } + }, + "disjointness_note": "All new ep0 ranges (300000-330099) are disjoint from every historical bank listed in sfm_protocol.py (12000, 20000+, 50000, 80000, 90000, 110000, 130000, 150000, 170000, 190000, 210000, 230000, 250000) and from the codex funnel banks (260000, 270000, 280000). Noise seeds 20260726-20260730 are new.", + "rule": "No episode or noise bank used to choose hyperparameters appears in final confirmation." +} diff --git a/overnight_run_07_12_sfm/claude_recipe_study/FIXED_RECIPE.json b/overnight_run_07_12_sfm/claude_recipe_study/FIXED_RECIPE.json new file mode 100644 index 0000000..540d778 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/FIXED_RECIPE.json @@ -0,0 +1,39 @@ +{ + "status": "FROZEN_BEFORE_FINAL_STUDY", + "declared_at": "2026-07-26T09:35:00-07:00", + "selected_by": "predeclared Stage-B qualification rule on the private M25 bank (ep0 310000, seed 20260728); winner = arm B2 best eligible round; no gathering-controller SR or training loss used", + "recipe": { + "name": "cost_hard_a0p01_e010_lr1em04_ess0p5", + "execution_selector": "safemppi_cost", + "replay_mode": "hard", + "replay_mode_semantics": "positives = population A only (certified windows with predicted min clearance <= 0.35 m AND window displacement >= 0.2 m); negatives = population B only (y=0 with actual collision, predicted window collision, declared trap rule, or displacement < 0.2 m); rules fixed in claude_offline_aug.py and declared before evaluation", + "alpha": 0.01, + "exposure_epochs": 10, + "lr": 1e-4, + "ess_target": 0.5, + "batch": 128, + "rounds": 10, + "seed": 20260724, + "K": 16, "B": 4, "T": 180, "H": 10, "nfe": 8, "temperature_gathering": 1.0, + "phi_s": 0.9, "gp_lambda": 0.01, "gp_cap": 512, + "scene_profile": "double_density_velocity_ood", + "expansion_bank_ep0": 20000, + "source_checkpoint_sha256": "1b5179c935d3eeff8824967d707d64cc9bab273949ee1f0e4f190172bab1b215", + "trainer": "claude_offline_exec_ext.py (reuses immutable gather/GP/beta core of sfm_b1_offline_exec.py)" + }, + "round_invariant": true, + "synthetic_recovery_mixture": "none in the frozen recipe (population C evaluated in Stage A/B: it repairs local NVP/collision contexts and raises Validity but does not preserve the liveness gate at any tested dose; B0-vs-B9 matched-round contrast showed no measurable marginal effect on top of margin/e100)", + "hard_data_mixture": "population A positives + population B negatives only, per replay_mode above", + "stage_b_evidence": { + "qual_bank_r0": {"SR": 0.669, "CR": 0.320, "Validity": 0.576, "clearance": 0.106, "time": 8.58}, + "winner_cell_B2_r1": {"SR": 0.714, "CR": 0.286, "Validity": 0.586, "clearance": 0.109, "time": 8.94}, + "eligible_runner_ups": [ + {"arm": "B3 cost/hard_recovery r1", "CR": 0.32}, + {"arm": "B8 cost/orig_plus_recovery ess.3 r1", "CR": 0.32}, + {"arm": "B1 cost/original r2", "CR": 0.33}, + {"arm": "B6 margin/hard_recovery lowdose r2", "CR": 0.34} + ], + "ineligible_high_validity_direction": "margin/original|orig_plus_recovery e100: V .69-.82, clearance up to .18, but SR fails the liveness gate at r1 and collapses to timeout from r2-r3 on the clean bank; reported as Pareto alternative, not selected" + }, + "checkpoint_selection": "SELECTION_RULE.json on the disjoint M50 bank (ep0 320000, seed 20260729) ONLY" +} diff --git a/overnight_run_07_12_sfm/claude_recipe_study/ITERATION2_PREREGISTRATION.json b/overnight_run_07_12_sfm/claude_recipe_study/ITERATION2_PREREGISTRATION.json new file mode 100644 index 0000000..c6de1f3 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/ITERATION2_PREREGISTRATION.json @@ -0,0 +1,33 @@ +{ + "status": "PREREGISTERED_BEFORE_READING_U10_U12", + "declared_at": "2026-07-26T13:15:00-07:00", + "motivation": "Iteration-1 frozen recipe (cost/hard) is a null result under its own predeclared rule (STAGE_D_SELECTION.json). The diagnostic companion cells on the M50 selection bank show B9-r1 (margin selector, original-plus-certified-recovery-v1 data, alpha .01, exposures 100, lr 1e-4, ESS .5) dominating r0 on CR/Validity/clearance while passing the liveness gate. Iteration 2 pre-registers that direction with fresh, unread banks.", + "candidate_recipes": { + "R_A": {"selector": "margin", "replay_mode": "orig_plus_recovery", "recovery_family": "v1", "alpha": 0.01, "exposure_epochs": 100, "lr": 1e-4, "ess_target": 0.5, "rounds": 4, "note": "checkpoints r1-r4 already exist (arm B9, trained from exact r0 BEFORE any selection-bank read)"}, + "R_B": {"selector": "margin", "replay_mode": "orig_plus_recovery_v2", "recovery_family": "v2 dodge-then-cruise (declared in claude_offline_aug.py)", "alpha": 0.01, "exposure_epochs": 100, "lr": 1e-4, "ess_target": 0.5, "rounds": 4} + }, + "freeze_criterion_declared_before_reading_U10_U12": "Compare single-update mirrors on the combined diag (ep0 300000, M8) + anchor (ep0 305000, M12) tuning banks: U10 = one r0 update with orig_plus_recovery_v2/e100/lr1e-4; U12 = identical with v1. Choose R_B iff U10 improves BOTH pooled SR and pooled successful time-to-goal versus U12 AND does not raise pooled CR by more than 0.01; otherwise choose R_A. Rationale: the v2 family exists to remove the conservative-escape slowdown; if it does not show that signature at matched dose on tuning banks, the extra novelty is not justified.", + "rounds_rationale": "rounds=4 is part of the recipe (a declared recipe variable): every e100 margin arm measured so far collapses to timeout from round 2-3; rounds beyond 4 are known-dominated and cost compute without adding selectable checkpoints.", + "selection": { + "bank": {"ep0": 340000, "noise_seed": 20260732, "m_per_gamma": 50, "role": "iteration-2 M50 checkpoint selection; NEVER read before this declaration"}, + "rule": "SELECTION_RULE.json applied verbatim (liveness gate SR >= SR(r0)-0.02, timeout <= timeout(r0)+0.05; then min CR, max Validity, max clearance, min time, smallest round) over r0..r4 of the chosen recipe" + }, + "confirmation": { + "bank": {"ep0": 330000, "noise_seed": 20260730, "m_per_gamma": 100, "role": "UNTOUCHED final confirmation: r0 + iteration-2 selected checkpoint + locked Kazuki; paired scenario-cluster CIs via claude_confirm_analysis.py"}, + "no_changes_after_reading": true + }, + "contamination_note": "The iteration-1 M50 bank (ep0 320000) was read for B9 r1/r2 as declared diagnostic companions and is therefore NOT used for iteration-2 selection. Banks 300000/305000/310000 are tuning banks. Bank 330000 has never been read by anything.", + "iteration_1_disposition": "reported in full as the rule-selected null result (frozen cost/hard recipe, winner r7), with its four-metric paper plot and all checkpoints; nothing about iteration 1 is hidden or reselected", + "freeze_criterion_applied": { + "applied_at": "2026-07-26T13:55:00-07:00", + "U_reads_combined_diag_anchor_n140": { + "r0": {"SR": 0.614, "CR": 0.386, "Validity": 0.571, "clearance": 0.129, "time": 8.19}, + "U12_v1_e100": {"SR": 0.614, "CR": 0.350, "Validity": 0.706, "clearance": 0.104, "time": 10.92}, + "U10_v2_e100": {"SR": 0.679, "CR": 0.264, "Validity": 0.723, "clearance": 0.109, "time": 11.22}, + "U11_v2_e10": {"SR": 0.564, "CR": 0.393, "Validity": 0.663, "clearance": 0.109, "time": 10.17} + }, + "evaluation": "U10 vs U12: SR improved (+.065) and CR improved (-.086), but successful time did NOT improve (+0.30 s). The declared criterion required BOTH SR and time improvements; v2 did not show its designed no-slowdown signature. Individual deltas are also within 95% noise at n=140.", + "decision": "R_A frozen (B9 knobs; existing checkpoints r1-r4 trained from exact r0 before any selection-bank read)", + "v2_disposition": "documented as a promising follow-up (goal-directed certified escapes raised single-update SR/CR at matched dose on tuning banks) — not selected, not confirmed, not claimed" + } +} diff --git a/overnight_run_07_12_sfm/claude_recipe_study/ITERATION2_SELECTION.json b/overnight_run_07_12_sfm/claude_recipe_study/ITERATION2_SELECTION.json new file mode 100644 index 0000000..280e392 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/ITERATION2_SELECTION.json @@ -0,0 +1,22 @@ +{ + "status": "ITERATION2_SELECTION_APPLIED", + "applied_at": "2026-07-26T15:05:00-07:00", + "recipe": "R_A frozen per ITERATION2_PREREGISTRATION.json: margin selector, orig_plus_recovery (family v1), alpha 0.01, exposure_epochs 100, lr 1e-4, ess_target 0.5, rounds 4, seed 20260724, expansion bank ep 20000+", + "bank": {"ep0": 340000, "noise_seed": 20260732, "m_per_gamma": 50}, + "rule": "SELECTION_RULE.json applied verbatim", + "per_round_pooled": { + "r0": {"SR": 0.563, "CR": 0.434, "timeout": 0.003, "Validity": 0.574, "clearance": 0.117, "time": 8.46}, + "r1": {"SR": 0.654, "CR": 0.311, "timeout": 0.034, "Validity": 0.715, "clearance": 0.117, "time": 10.82}, + "r2": {"SR": 0.537, "CR": 0.300, "timeout": 0.163, "Validity": 0.789, "clearance": 0.132, "time": 13.17}, + "r3": {"SR": 0.020, "CR": 0.329, "timeout": 0.651, "Validity": 0.816, "clearance": 0.209, "time": 14.46}, + "r4": {"SR": 0.000, "CR": 0.331, "timeout": 0.669, "Validity": 0.821, "clearance": null, "time": null} + }, + "liveness_gate": {"SR_min": 0.543, "timeout_max": 0.053}, + "eligible_rounds": ["r1"], + "selected_round": "r1", + "selected_checkpoint": "/data3/research1/claude_sfm_best_recipe_f06e8dd/stageB/B9_margin_origrec_e100/round_01.pt", + "selected_vs_r0_on_selection_bank": {"dSR": 0.091, "dCR": -0.123, "dValidity": 0.141, "dclearance": 0.000, "dtime": 2.36}, + "stability_note": "r2 fails the gate (timeout .163) and r3-r4 collapse: the recipe's gains are a one-round phenomenon; this is reported plainly (not claimed as monotonic learning).", + "paper_plot": "/data3/research1/claude_sfm_best_recipe_f06e8dd/iteration2/paper_trends/b1_margin50_metric_trends.{png,pdf} + numeric SEs in margin_origrec_e100_trends_rows.jsonl", + "next": "M100 confirmation on untouched bank 330000 (r0 + selected r1 + locked Kazuki); no changes after reading" +} diff --git a/overnight_run_07_12_sfm/claude_recipe_study/MPC_STUDY_DELIVERY.json b/overnight_run_07_12_sfm/claude_recipe_study/MPC_STUDY_DELIVERY.json new file mode 100644 index 0000000..7ba1e1c --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/MPC_STUDY_DELIVERY.json @@ -0,0 +1,38 @@ +{ + "status": "MPC_DISTILLATION_STUDY_COMPLETE_FAIL_CLOSED_AT_M10", + "completed_at": "2026-07-26T21:35:00-07:00", + "branch": "agent/claude-sfm-mpc-distill-20260726 (Codex branch @0e0eca2 untouched, read-only)", + "primary_question": "does dedicated D_MPC+ distillation produce a VISIBLE raw-policy change, rather than merely lowering its training loss?", + "answer": "YES — visible and consistent, but not globally beneficial at any swept dose.", + "evidence_visible_change": { + "socp_positive_rate_at_mpc_contexts": "rose after 15 of 16 dedicated blocks (e.g. .106->.457 at lr 1e-4 x8 epochs, .106->.482 at lr 3e-4 x2; only M4 round-4 fell .319->.278), measured by 16 fresh temperature-1 raw samples per held gamma-balanced MPC context against the exact full-H10 SOCP verifier", + "target_recovery_rmse": "fell after 16 of 16 blocks (policy samples move toward the stored MPC targets)", + "loss_vs_policy": "distill_block_loss_vs_policy_change.json: policy-side movement occurs even in blocks where the CFM loss barely moved or rose (M1 r1: loss 1.248->1.266 while SOCP rate .106->.216), so the change is not a loss artifact" + }, + "evidence_not_beneficial": { + "m10_selection": "MPC_M10_SELECTION.json: 0 of 32 pre/post cells pass the pre-registered liveness gate (SR >= r0-0.02 = .637, timeout <= .079); r0 on the bank: SR .657 / CR .314 / V .628", + "immediate_liveness_cost": "round-1 post-block SR: .629 -> .471 (lr1e-4x2), .257 (lr1e-4x8), .243 (lr3e-4x2), .600 (lr1e-5x8 — gentlest, also smallest local gain)", + "no_collapse_rescue": "the base ordinary recipe (margin/original/e10) collapses by rounds 2-3 with or without distillation (control_no_distill: SR .686 -> .457 -> .000)", + "funnel": "per the pre-registration, the M50 (360000) and M100 (370000) confirmations are NOT triggered because M10 selects no eligible winner; stopping fail-closed rather than promoting an ineligible cell" + }, + "dedicated_recipe_sweep_reported": { + "M1": {"lr_dedicated": 1e-4, "epochs": 2}, "M2": {"lr_dedicated": 1e-4, "epochs": 8}, + "M3": {"lr_dedicated": 1e-5, "epochs": 8}, "M4": {"lr_dedicated": 3e-4, "epochs": 2}, + "rounds": 4, "base_ordinary_recipe": "margin selector, original replay, alpha .01, exposures 10, lr 1e-4, ESS .5, seed 20260724", + "checkpoints": "every pre-block and post-block checkpoint kept under mpc_distill/M*/", + "exact_choice": "none selected — no cell eligible under the pre-registered rule" + }, + "d_mpc_plus_accounting": { + "kept_per_round_range": [61, 222], + "intersection_semantics": "privileged exact-SFM look-ahead feasible AND exact full-H10 SOCP y=1, ranked by frozen native SafeMPPI cost, top-2 per hard context", + "separation_verified": "records live in per-round d_mpc_plus_round_*.pt only; ordinary D/D+/GP/acquisition untouched (code path in claude_mpc_exec.py); prefix-replay reconstruction fail-closed (0 mismatches in all arms)" + }, + "synthesis_with_recipe_study": "Across three target families (deterministic v1 escapes, goal-directed v2, privileged Codex MPC pool), concentrated hard-context-only distillation always trades global liveness for local certifiable support. The confirmed winning recipe embedded certified hard-context targets INSIDE the full original replay mixture (~5% mass) — composition, not target quality, is the binding constraint.", + "artifacts": { + "selection": "mpc_distill/MPC_M10_SELECTION.json (all 37 cells incl. control and r0)", + "audits": "mpc_distill/M*/metrics.jsonl (per-block before/after local + M10 audits)", + "loss_vs_change": "mpc_distill/distill_block_loss_vs_policy_change.json", + "buffers": "mpc_distill/M*/d_mpc_plus_round_*.pt (with per-record provenance and verifier diagnostics)", + "control": "mpc_distill/m10_r0_and_control/" + } +} diff --git a/overnight_run_07_12_sfm/claude_recipe_study/MPC_STUDY_PREREGISTRATION.json b/overnight_run_07_12_sfm/claude_recipe_study/MPC_STUDY_PREREGISTRATION.json new file mode 100644 index 0000000..5864e21 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/MPC_STUDY_PREREGISTRATION.json @@ -0,0 +1,27 @@ +{ + "status": "PREDECLARED_BEFORE_ANY_MPC_ARM_RUN", + "declared_at": "2026-07-26T16:20:00-07:00", + "branch": "agent/claude-sfm-mpc-distill-20260726 (Codex branch agent/sfm-adhoc-controller-overlay-20260726@0e0eca2 kept read-only; nothing written to it)", + "pipeline": "claude_mpc_pool.py + claude_mpc_exec.py: ordinary B1 expansion round (margin selector, original replay, alpha .01, exposures 10, lr 1e-4, ESS .5) followed by a dedicated D_MPC+ distillation block per round", + "d_mpc_plus": { + "separation": "per-round buffer; never enters D/D+, GP buffer, or acquisition; privileged controller never used at evaluation; no x0 stored (cfm_loss draws fresh Gaussian bases)", + "hard_contexts": "population-B windows (declared rules in claude_offline_aug.py) UNION all NVP contexts of the round shard; capped at 400/round by even stride", + "pool": "original Codex privileged pool regenerated at the stored context after exact prefix replay of the live SFM crowd (fail-closed on reconstruction mismatch): guided flow (n_sample 200) + MPPI refinement + brake + 12 goal plans + 18 avoidance plans + 25 constant-acceleration escapes", + "filter": "privileged exact-SFM look-ahead feasible (recoverable AND clearance >= per-gamma hard margin .22-.24) AND exact full-H10 SOCP y=1; intersection ranked by frozen native SafeMPPI proposal cost; keep top 2 per context" + }, + "sweep_arms_rounds_4_each": { + "M1": {"lr_dedicated": 1e-4, "epochs_dedicated": 2}, + "M2": {"lr_dedicated": 1e-4, "epochs_dedicated": 8}, + "M3": {"lr_dedicated": 1e-5, "epochs_dedicated": 8}, + "M4": {"lr_dedicated": 3e-4, "epochs_dedicated": 2} + }, + "control": "the identical ordinary recipe WITHOUT distillation = codex arm offline_exec_alpha0p01_exposures010 (same seeds, rounds 1-4 checkpoints already archived); its cells are evaluated on the same M10 bank for the comparison", + "banks": { + "m10_audit_and_selection": {"ep0": 350000, "noise_seed": 20260733, "m_per_gamma": 10, "role": "before/after audit of every dedicated block AND checkpoint/arm selection"}, + "m50_confirmation": {"ep0": 360000, "noise_seed": 20260734, "m_per_gamma": 50, "role": "disjoint confirmation of the single M10-selected winner vs r0"}, + "m100_final": {"ep0": 370000, "noise_seed": 20260735, "m_per_gamma": 100, "role": "disjoint final confirmation of the M50-confirmed winner vs r0; only if M50 confirms a real effect"} + }, + "selection_rule": "on the M10 bank over all saved checkpoints (pre-block and post-block, every arm): liveness gate SR >= SR(r0)-0.02 and timeout <= timeout(r0)+0.05; among eligible minimize CR, then maximize Validity, then clearance, then minimize time; the winner must be a POST-BLOCK checkpoint to count as a distillation effect, and its pre-block sibling is always reported alongside for the marginal-attribution claim", + "primary_question": "does the dedicated D_MPC+ block produce a VISIBLE raw-policy change (SOCP-positive rate / progress / clearance / target recovery at MPC contexts, and fixed-bank M10 deltas) beyond lowering its own training loss?", + "prohibitions": ["no temperature tuning", "no test-bank tuning", "no hidden fallback", "no silent rollback (every pre/post checkpoint kept and reported)", "no modification of any Codex artifact"] +} diff --git a/overnight_run_07_12_sfm/claude_recipe_study/R1_CONTINUATION_DELIVERY.json b/overnight_run_07_12_sfm/claude_recipe_study/R1_CONTINUATION_DELIVERY.json new file mode 100644 index 0000000..ccb216c --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/R1_CONTINUATION_DELIVERY.json @@ -0,0 +1,44 @@ +{ + "status": "R1_STABLE_CONTINUATION_DELIVERY_COMPLETE", + "completed_at": "2026-07-27T05:40:00-07:00", + "branch": "agent/claude-sfm-r1-stable-continuation-20260727", + "headline": "No stable multi-round improvement exists under the pre-registered hard gates: 0/16 round-1 candidates admissible (timeout gate binds at every dose, including E=1/lr=1e-5/drift 0.0018). Frozen winner = immutable r1. Final M100 verdict: NOT a win over Kazuki — Validity decisively ahead (+.384 [+.354,+.412], more than double), CR (+.089), clearance (-.055) and time (+6.49 s) remain behind.", + "development": { + "rounds_run": 1, "rounds_accepted": 0, + "all_candidate_rows": "r1_continuation/develop/metrics.jsonl (16 combos x replay logs, route diversity, drift, admissibility checks; nothing hidden)", + "diagnostics_figure": "r1_continuation/diagnostics/round1_candidates.png", + "key_numbers": { + "timeout_gate": 0.093, "candidate_timeouts": [0.129, 0.229], + "best_CR_seen": 0.129, "best_Validity_seen": 0.789, + "route_entropy_range": [0.973, 1.008], + "representation_cosine_range": [0.9879, 0.9982] + }, + "anchor_buffer": "r1_continuation/r1_anchor_buffer.pt (1400 records, 200/gamma, all exact-certified, from 57/84 successful r1 raw rollouts on bank 380000)" + }, + "final_m100": { + "bank": {"ep0": 410000, "noise_seed": 20260742, "m_per_gamma": 100, "first_read": "after R1_CONTINUATION_SELECTION.json was committed"}, + "pooled": { + "r0": {"SR": 0.696, "CR": 0.299, "timeout": 0.006, "Validity": 0.619, "clearance": 0.119, "time": 8.62}, + "r1": {"SR": 0.690, "CR": 0.277, "timeout": 0.033, "Validity": 0.736, "clearance": 0.129, "time": 10.81}, + "kazuki": {"SR": 0.810, "CR": 0.189, "timeout": 0.001, "Validity": 0.353, "clearance": 0.183, "time": 4.32} + }, + "paired_r1_minus_kazuki": { + "CR": [0.0886, 0.0186, 0.1571], "validity": [0.3839, 0.3542, 0.4122], + "clearance": [-0.0547, -0.0705, -0.0389], "time": [6.4910, 6.1666, 6.8141] + }, + "paired_r1_minus_r0": { + "CR": [-0.0214, -0.0714, 0.0257], "validity": [0.1177, 0.0997, 0.1358], + "clearance": [0.0097, -0.0030, 0.0224], "time": [2.1974, 1.9731, 2.4241] + }, + "gamma_trend_r1": "gamma=0.1: clearance .144 (largest), time 13.79 s (longest) — ordering intact", + "verdict": "FULL PARETO DOMINANCE OVER KAZUKI NOT ACHIEVED; metrics behind: CR, successful clearance, successful time. Metric ahead: window Validity (raw flow certificate-conformance is 2.1x Kazuki's). Reported exactly; not relabeled." + }, + "mechanism_synthesis": "Every continuation gradient step along the safety axis pays in time-to-goal at T=180; the self-anchor (mass .50 > .25 on SR) slows but does not stop the trade. Kazuki's advantage is runtime computation (goal/CBF guidance + MPPI refinement per step), which the raw temperature-1 flow cannot amortize into weights on this OOD profile without losing liveness. The raw policy's one decisive edge is exact-certificate Validity.", + "artifacts": { + "final_metrics": "r1_continuation/final_m100/ + r1_continuation/final_kazuki_m100/", + "paired_cis": "r1_continuation/FINAL_CONFIRMATION_ANALYSIS.json", + "four_metric_curves": "r1_continuation/final_m100/raw_m100_offline_curves.{png,pdf}", + "prereg_chain": ["R1_CONTINUATION_PREREGISTRATION.json", "R1_CONTINUATION_SELECTION.json (frozen pre-read)"], + "checkpoints": "r1 (immutable) + all 16 round-1 candidate checkpoints under r1_continuation/develop/round_01/" + } +} diff --git a/overnight_run_07_12_sfm/claude_recipe_study/R1_CONTINUATION_PREREGISTRATION.json b/overnight_run_07_12_sfm/claude_recipe_study/R1_CONTINUATION_PREREGISTRATION.json new file mode 100644 index 0000000..626123d --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/R1_CONTINUATION_PREREGISTRATION.json @@ -0,0 +1,52 @@ +{ + "status": "PREDECLARED_BEFORE_ANY_CONTINUATION_RUN", + "declared_at": "2026-07-27T00:15:00-07:00", + "branch": "agent/claude-sfm-r1-stable-continuation-20260727 (all prior branches/checkpoints/artifacts preserved untouched)", + "start_checkpoint": { + "path": "/data3/research1/claude_sfm_best_recipe_f06e8dd/stageB/B9_margin_origrec_e100/round_01.pt", + "sha256_prefix": "141e4ae6592bf73f", + "immutable": true + }, + "prohibitions": ["no e100 continuation", "no MPC distillation", "no Kazuki retuning", "no temperature tuning", "final M100 bank (ep0 410000) never read before the single winner is frozen"], + "fixed_mechanism": { + "selector": "margin", "replay": "original D+ and D- of the new shard, plus exact-SOCP-certified deterministic recovery positives (family v1), plus the frozen r1 self-anchor buffer", + "recovery_positive_mass_share": 0.05, + "alpha": 0.01, "ess_target": 0.5, + "gp_verifier_scene_gammas_encoder_gathering_eval": "identical to the confirmed study; frozen encoder SHA-checked; raw evaluation temperature 1.0 NFE 8", + "free_variables_only": ["effective exposures E in {1,4,10,25}", "lr in {1e-5, 3e-5}", "r1-anchor positive mass share in {0.25, 0.50}"] + }, + "positive_mass_composition": "anchor_mass*anchor + 0.05*recovery + (0.95 - anchor_mass)*new-shard D+; within each population the standard hierarchy mass (gamma -> episode -> context -> record); negatives = new-shard D- with the unchanged alpha=.01 signed-gradient scheme", + "anchor_buffer": { + "definition": "r1 raw temperature-1 rollouts on the private anchor bank; only SUCCESSFUL episodes; only executed sliding windows with full H_t=10 that are exact-verifier y=1; contexts recorded exactly as the gathering store (hp10/low5/hist/state/ped_xy/ped_vel)", + "bank": {"ep0": 380000, "m_per_gamma": 12, "noise_seed": 20260739}, + "cap": "200 windows per gamma, deterministic selection: order by (episode, step), take first 200", + "semantics": "self-replay of the accepted policy's own certified successful behavior; not expert data; frozen once before round 1 and never regenerated" + }, + "continuation_procedure": { + "rounds_max": 10, + "gather": "one shard per round from the CURRENTLY ACCEPTED checkpoint; continuation round k uses expansion_scenarios(1+k) (scenarios 20008..20087), preserving the original bank progression; GP previous-shard chain starts from the archived B9 round-1 shard; ell = 0.3259987743470518 (B9 manifest)", + "shared_across_candidates": "identical shard, identical recovery records, identical replay seed and deterministic batch ordering; only (E, lr, anchor_mass) differ; every candidate forks from the accepted checkpoint", + "candidates_per_round": 16, + "logged_per_candidate": ["unique sample exposure", "gradient norms (positive/negative)", "parameter drift per module", "penultimate-representation drift on a fixed probe batch", "new-shard positive loss", "anchor loss", "recovery loss", "route diversity (mode distribution + entropy of 16 raw samples at 32 fixed gamma-balanced probe contexts)"], + "rejected_candidates": "all 16 rows published every round; nothing hidden" + }, + "development_bank": {"ep0": 390000, "noise_seed": 20260740, "m_per_gamma": 10, "role": "fixed raw M10/gamma development screening; r1 evaluated once on it as the admissibility baseline"}, + "admissibility_vs_r1_dev_values": { + "SR": "SR >= SR_r1 - 0.05 (predeclared tolerance)", + "timeout": "timeout <= timeout_r1 + 0.05 (predeclared tolerance)", + "CR": "CR <= CR_r1 (hard, no tolerance)", + "Validity": "V >= V_r1 (hard)", + "clearance": "successful clearance >= clr_r1 (hard)", + "gamma_trend": "clearance(gamma=0.1) >= clearance(gamma=1.0) AND successful time(gamma=0.1) >= successful time(gamma=1.0); a gamma cell with zero successes fails the check", + "note": "hard inequalities exactly as directed; if no candidate is admissible the procedure stops (no forced rounds)" + }, + "selection_among_admissible": "lexicographic: lowest CR, then highest Validity, then highest clearance, then shortest successful time; winner becomes the accepted checkpoint", + "fixed_recipe_choice_rule": "after development stops: the (E, lr, anchor_mass) combo accepted most often; ties broken by better mean lexicographic rank across its accepted rounds, then by smaller E, then smaller lr, then smaller anchor_mass; rerun ALL continuation rounds from immutable r1 with that single combo (no round-dependent tuning); stability requirement = at least three consecutive rerun checkpoints admissible on the development bank", + "funnel": { + "screening": "fixed raw M10/gamma (development bank) for every rerun round", + "m50": {"ep0": 400000, "noise_seed": 20260741, "m_per_gamma": 50, "role": "fresh disjoint eligibility for admissible rerun checkpoints"}, + "m100": {"ep0": 410000, "noise_seed": 20260742, "m_per_gamma": 100, "role": "UNTOUCHED final: exactly one frozen winner vs r0, r1, and locked Kazuki, paired scenario-cluster CIs"} + }, + "win_criterion": "full Pareto dominance over locked Kazuki on final M100: lower CR AND higher window Validity AND higher successful clearance AND shorter successful time, with SR/liveness preserved and the gamma ordering intact; anything less is reported metric-by-metric and NOT relabeled a win", + "compute": "GPU3 only (GPU1 now occupied by another user; exclusivity rule); exact SOCP verification remains CPU" +} diff --git a/overnight_run_07_12_sfm/claude_recipe_study/R1_CONTINUATION_SELECTION.json b/overnight_run_07_12_sfm/claude_recipe_study/R1_CONTINUATION_SELECTION.json new file mode 100644 index 0000000..d7846a4 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/R1_CONTINUATION_SELECTION.json @@ -0,0 +1,27 @@ +{ + "status": "R1_CONTINUATION_SELECTION_FROZEN_BEFORE_FINAL_M100_READ", + "declared_at": "2026-07-27T03:05:00-07:00", + "development_outcome": { + "rounds_run": 1, + "rounds_accepted": 0, + "stop_reason": "pre-registered rule: no admissible candidate at round 1; procedure stops, no forced rounds", + "dominant_failure": "TIMEOUT gate (candidates .129-.229 vs gate .093): every continuation dose, including the smallest (E=1, lr=1e-5, trunk drift 0.0018), immediately trades liveness for safety at the T=180 horizon", + "secondary_observations": [ + "CR improves monotonically with dose down to .129 and Validity to .789 (E25/lr1e-5/anchor.50 fails ONLY the timeout gate)", + "anchor mass .50 consistently preserves more SR than .25 (goal-seeking anchor works directionally)", + "all 16 candidate rows, replay logs, route-diversity and drift diagnostics published in r1_continuation/develop/metrics.jsonl; nothing hidden" + ] + }, + "fixed_recipe": "none exists — zero accepted rounds means there is no (E, lr, anchor_mass) continuation recipe to confirm; the fixed-recipe rerun degenerates to zero rounds", + "frozen_winner": { + "checkpoint": "/data3/research1/claude_sfm_best_recipe_f06e8dd/stageB/B9_margin_origrec_e100/round_01.pt", + "sha256_prefix": "141e4ae6592bf73f", + "identity": "the immutable r1 selected and confirmed in the prior study; unchanged" + }, + "final_m100": { + "bank": {"ep0": 410000, "noise_seed": 20260742, "m_per_gamma": 100}, + "methods": ["r0", "r1 (frozen winner)", "locked Kazuki (.3/.5)"], + "read_only_after_this_declaration": true, + "verdict_rule": "full Pareto dominance over Kazuki (lower CR AND higher Validity AND higher clearance AND shorter time, SR/liveness preserved, gamma ordering intact) or a metric-by-metric behind-report; no relabeling" + } +} diff --git a/overnight_run_07_12_sfm/claude_recipe_study/RESEARCH_LOG.md b/overnight_run_07_12_sfm/claude_recipe_study/RESEARCH_LOG.md new file mode 100644 index 0000000..69a90c7 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/RESEARCH_LOG.md @@ -0,0 +1,119 @@ +# Claude SFM Safe Flow Expansion — best fixed recipe study + +- Private worktree: `/home/dohyun/projects/safeMPPI-claude-sfm-recipe-f06e8dd` +- Branch: `agent/claude-sfm-best-recipe-20260726` from `f06e8ddc11fc7a1f5bada2cb2a587bff3ba4e424` (origin/agent/sfm-b1-offline-eval-funnel) +- Output root: `/data3/research1/claude_sfm_best_recipe_f06e8dd` +- r0 checkpoint: `/home/dohyun/projects/sfm_hp10_b1_runs/103476d/pretrained_hp10.pt` + - SHA-256 verified `1b5179c935d3eeff8824967d707d64cc9bab273949ee1f0e4f190172bab1b215` (matches `safe_flow_expansion_SFM@491e478:checkpoints/hp10_pretrained_r0.pt` blob hash) +- Plotting contract checkout: `/home/dohyun/projects/safe_flow_expansion-claude-plot-87063d3` @ 87063d3 (read-only) +- GPUs: physical 1 and 3 (verified idle at claim time: 18–19 MiB, 0% util). Exact SOCP verifier stays CPU. +- Python: conda env `cfm_mppi` (Python 3.11). + +## 2026-07-26 02:30 — starting state inherited from Codex + +Prior completed work (all on `double_density_velocity_ood`, 40 peds, 1.0–2.0 m/s): + +- 27 trained arms (3 selectors × α∈{0,.01,.1} × exposure∈{1,10,100}), rounds 0–10, lr 1e-4, ESS 0.5, batch 128, seed 20260724, expansion bank ep 20000+: + - margin: `/data3/research1/sfm_b1_offline_exec_9arm_f36393b` + - safemppi_cost: `/data3/research1/sfm_b1_offline_cost_9arm_db80dfc` + - balanced_rank: `/data3/research1/sfm_b1_offline_balanced_9arm_97630c6` +- Funnel + final M100 (completed 2026-07-26 05:49 UTC): `/data3/research1/sfm_b1_offline_final_funnel_f06e8dd` + - M100 (ep0 280000, seed 20260725): r0 SR .660 CR .337 V .595 clr .118 t 8.65 + - global winner (cost α.01 exp10 r8): SR .690 CR .310 V .531 clr .117 t 7.52 — CR −2.7pt but Validity −6.4pt: **not** a genuine fixed-recipe win. + +### Key measured phenomenon driving this study + +Per-round fixed-raw M50 curves (margin 9-arm factorial CSV) show **SR collapse to 0.00 by round 3–4 in every margin arm** while Validity rises to ~.85 — a monotone slowdown (successful time 8.7→10.7→13→16.4 s) until nothing reaches the goal inside T=180. Balanced_rank arms collapse identically (M10 screening). Cost-selector arms avoid the collapse but mostly degrade (CR rises to .4–.6 in 6/9 arms). Only round-1/2 checkpoints ever beat r0, modestly (best M50: margin α.01 exp100 r1 = SR .694, CR .28, V .74, clr .128, t 10.7). + +Working hypothesis (to be tested in Stage A): replay trains the flow toward the *gathering controller's* executed-window distribution, which is verifier-gated and conservative; each round compounds the slowdown. The failure is data composition, not just step size. + +## Plan + +Stages A–E per task spec; banks in `EPISODE_BANKS.json`. Interventions implemented as opt-in modules, default OFF, original behavior preserved with existing tests. + +## 2026-07-26 03:20 — modules, tests, baselines launched + +- New additive modules (commit 779d2de): `claude_offline_aug.py` (declared pop-A/B rules + deterministic certified-recovery generator; family = 145 candidates/context, prefilter cap 24, keep ≤2, exact `SM.verify_query` gate), `claude_offline_exec_ext.py` (opt-in lr/ESS/rounds/replay-mode; immutable core reused), `claude_stageA.py` (mine/branch/update/compare), `claude_kazuki_eval.py` (locked comparator + executed-window Validity), `claude_paper_trends.py` (evaluator → paper plot contract). +- Existing tests: 43 passed. New tests: 6 passed, incl. bitwise default-OFF equivalence of `replay_with_mode("original")` vs `OR.replay`, no-relabel guarantee, and independent exact recertification of every synthetic recovery record. +- Baselines launched on the pre-registered codex M100 bank (ep0 280000, seed 20260725): raw r0 + margin winner (α.01 exp100 r1) + cost winner (α.01 exp10 r8) on GPU1; locked Kazuki on GPU3. +- Stage A mining (codex round-1 shards): margin arm — 1034/4014 NVP contexts, 22 collision windows, 310 trap windows; cost arm — 1102 NVP, 19 collisions, only 32 traps. D+ plan-geometry drift r1→r3 (median displacement 0.95→0.70 m) supports the composition-drift hypothesis. +- Stage A single-update candidates U1–U8 launched (one replay round on r0 from archived round-1 shards; margin + cost shards × {control, lr1e-5, exp1, hard, hard_recovery, lowdose-hardrec}). + +## PREDECLARED Stage-B qualification rule (written before any Stage-B run) + +- Bank: qualification raw bank ep0 310000, noise seed 20260728, M=25/γ, temperature 1.0, via `sfm_b1_offline_eval.py` only. No gathering-controller SR, no training loss. +- Every candidate arm trains rounds 1–4 from the exact r0 checkpoint (seed 20260724, expansion bank ep 20000+, identical gather semantics). +- Eligibility per round r ∈ {1..4}: SR(r) ≥ SR(r0) − 0.02 AND timeout(r) ≤ timeout(r0) + 0.05 on the qualification bank (liveness gate; r0 evaluated on the same bank/noise). +- Arm score = its best eligible round ordered by (min CR, then max Validity, then max successful clearance, then min successful time-to-goal). Arms with no eligible round are disqualified (collapse). +- Stability tie-break: among arms whose best-round CR are within 0.03 of the leader, prefer the arm whose round-4 checkpoint is still eligible; among those, the better round-4 CR. Rationale: the frozen Stage-C/D recipe must hold 10 round-invariant macro-rounds. +- The frozen recipe = the winning arm's knobs verbatim; final study rounds fixed at 10; Stage-D checkpoint selection governed solely by SELECTION_RULE.json on the disjoint M50 bank (ep0 320000). + +## 2026-07-26 05:30 — Stage A results (complete) + +**Mechanism (r0 branch trace, 24 gathering lineages, all-K exact verification):** 2117 contexts, 456 NVP (21.5%); only 49 (2.3%) had a positive in K that B missed → the flow itself lacks certifiable support at hard contexts; B=4 is not binding. NVP concentrates at episode start (43–51% in steps 0–39 → ~0 after step 60): origin-corner congestion with 40 fast pedestrians. Every collision episode dies through a terminal run of K+=0 NVP contexts with negative predicted clearance (uncertified raw execution). Selector disagreement at 59% of contexts. Certified deterministic escapes exist at ~53% of hard contexts (269/506 margin shard, 202/501 cost shard). + +**Single-update raw reads (diag M8 + anchor M12, n=140/ckpt):** r0 SR .614 / CR .386 / V .571. +- U6 cost-shard original: SR .714 / CR .286 / V .566, faster — only arm improving SR/CR/time; Validity flat. +- U1–U5 margin-shard arms: Validity +.08–.13 (U5 hard+recovery best, .698) but SR −.03–.07 and slower — conservative drift visible after ONE update. +- U7 cost-shard hard-only: WORSE than U6 across the board — discarding the goal-directed mass hurts. +- U8 lowdose hardrec: mild moves, dose too small per round. + +**Local repair (before/after branch traces, identical keyed latents):** U5 hard+recovery cut NVP 456→341 (−25%), raised B-positive fraction .712→.844, repaired the three hardest collision lineages (s20005 γ.1/γ.5, s20004 γ.5 → success), introduced slowdown timeouts elsewhere. U6 sped the policy up but *lowered* gathering certifiability (B+ frac .592, NVP 483). Conclusion: recovery data provides certifiable support exactly where the flow lacks it; the composition question (keep goal-seeking mass + add recovery) is what Stage B arms B4/B5/B8 test (`orig_plus_recovery`, declared before evaluation). + +**Kazuki locked baseline (M100 ep0 280000):** SR .779 / CR .217 / **Validity .350** / clearance .181 / time 4.36 s — fast and lower-CR than r0 but far below r0 on exact-certificate Validity (.35 vs .59). + +## Baseline reproduction complete (M100, ep0 280000, seed 20260725) + +| method | SR | CR | timeout | Validity | succ. clearance | succ. time | +|---|---:|---:|---:|---:|---:|---:| +| r0 raw (repro) | .6586 | .3386 | .0029 | .5944 | .1176 | 8.664 | +| r0 raw (codex funnel) | .6600 | .3371 | .0029 | .5945 | .1179 | 8.646 | +| B1 control: margin α.01 e100 **r1** | .7314 | .2371 | .0314 | .7390 | .1336 | 10.733 | +| B1 control: cost α.01 e10 **r8** (codex global winner) | .6900 | .3100 | .0000 | .5311 | .1169 | 7.520 | +| locked Kazuki (.3/.5) | .7786 | .2171 | .0043 | .3495 | .1809 | 4.355 | + +r0 reproduces codex within 1–2 flipped episodes (GPU FP nondeterminism). The margin-r1 control dominates the codex-selected cost-r8 on this bank — but it is a pre-collapse snapshot (SR→0 by r3–4 in that arm). Bar for the new fixed recipe: margin-r1-level CR/Validity gains with multi-round stability. + +## Stage B launched 05:20 — 8 arms × 4 rounds from exact r0 + +B1 cost/original, B2 cost/hard, B3 cost/hard_recovery, B4 cost/orig_plus_recovery (all α.01 e10 lr1e-4 ess.5); B5 cost/orig_plus_recovery lowdose (e1 lr1e-5); B6 margin/hard_recovery lowdose; B7 cost/original ess.3; B8 cost/orig_plus_recovery ess.3. Qualification: predeclared rule on M25 bank ep0 310000 (see above). + +### Stage B qualification (M25, ep0 310000; r0 = SR .669 / CR .320 / V .576 / clr .106 / t 8.58) + +Per arm r1..r4 (SR/CR/V): +- B1 cost orig: .60/.39/.54, .67/.33/.53, .64/.35/.53, .64/.36/.51 — no gain, V drifts down +- B2 cost hard: .71/.29/.59 then degrades to .62/.38/.52 +- B3 cost hardrec: .68/.32/.60 then degrades to .55/.45/.55 +- B4 cost origrec: ≈flat (.60–.66 SR, V .57–.59) +- B5 cost origrec lowdose: flat +- B6 margin hardrec lowdose: V .58→.65 climbing, CR .34–.40 (no CR gain), t 8.9→10.2 — slow conservative drift +- B7 cost orig ess.3: worse than B1 (lower-ESS acquisition does not help) +- B8 cost origrec ess.3: ≈B4 +Verdict: no cost-selector composition materially improves CR or Validity on this bank; the lower ESS target (0.3) is not beneficial. The strongest known pattern (margin/original/e100 — codex arm, r1 CR .237/V .739 on the M100 280k baseline) was absent from the set. + +## 2026-07-26 13:00–15:00 — Stage D result, mechanism deep-dive, iteration 2 + +**Stage D (frozen cost/hard recipe, M50 bank 320000): honest null.** Rule-selected r7: CR .263 vs r0 .269 but Validity .571 vs .638. Full 11-round curve + paper-contract plot at `stageD/paper_trends/`. Recorded in STAGE_D_SELECTION.json with Pareto frontier. + +**Mechanism figures** (user request; artifact https://claude.ai/code/artifact/3c16fe2a-1450-4a8c-addf-466beaea1111, PNGs in `~/claude_sfm_figs/` and `mechanism/`): the closing certifiability window quantified on stored contexts — onset−2: 10–13/16 flow candidates certify, escapes 14–24/24; onset: flow 1–2/16, escapes 5–15/24; onset+4–7: 0 everywhere including both deterministic families. Closed-loop replays: B9-r1 converts both collision lineages to successes (s20005 γ.1 clearance .236; s20004 γ.5 clearance .139). + +**Recovery family v2** (user hypothesis: v1 escapes too conservative): dodge-then-cruise family implemented + declared; single-update head-to-head at matched dose (n=140): U10 v2 SR .679/CR .264/V .723/t 11.22 vs U12 v1 SR .614/CR .350/V .706/t 10.92. v2 did NOT remove the slowdown (the declared signature), so per the pre-registered criterion iteration 2 froze R_A (v1); the v2 SR/CR edge (within noise) is documented as follow-up. + +**Iteration 2 (pre-registered):** recipe = margin/orig_plus_recovery-v1/α.01/e100/lr1e-4/ess.5/rounds4 (B9; checkpoints trained from exact r0 before any selection-bank read). Fresh M50 selection bank 340000: r0 = SR .563/CR .434/V .574; **r1 = only eligible round: SR .654 (+.091), CR .311 (−.123), V .715 (+.141), clearance flat, time +2.36 s**; r2 fails gate (timeout .163), r3–r4 collapse. Selected checkpoint: B9 round_01.pt. Population counts + per-row certificate audits in `iteration2/population_counts_B9.json` and `iteration2/recovery_certificate_audit_B9.json` (280–407 exact-certified recovery positives/round from 8.3–12.3k exact queries). + +**Stage E launched** on untouched M100 bank 330000: r0 + selected + locked Kazuki. + +## 2026-07-26 18:20 — FINAL: Stage E confirmation + delivery + +M100 confirmation (untouched bank 330000, 700 CRN rollouts/method): r0 SR .643/CR .353/V .603/clr .108/t 8.76 → **selected (margin+orig∪recovery-v1, e100, round 1): SR .723 / CR .253 / V .736 / clr .122 / t 10.84**; locked Kazuki SR .827/CR .173/**V .353**/clr .168/t 4.17. Paired scenario-cluster 95% CIs (selected − r0): ΔCR −.100 [−.157,−.043], ΔV +.133 [+.114,+.153], Δclr +.014 [+.005,+.023], Δt +2.08 [1.84,2.32] — all exclude zero. γ=0.1 keeps the largest clearance (.148) and longest time (13.4 s). Stability honestly reported: round-1 phenomenon; r2+ collapse. Full record in DELIVERY_COMPLETE.json; 33-artifact SHA manifest. + +## 2026-07-26 16:20+ — MPC distillation follow-up study (separate branch) + +Pre-registered in MPC_STUDY_PREREGISTRATION.json (banks M10 350000 / M50 360000 / M100 370000). Smoke round: 172 D_MPC+ records; dedicated block doubled the local SOCP-positive rate at MPC contexts (.161→.313) — a visible raw-policy change; M10 CR moved adversely at that dose (smoke bank). 4-arm sweep (lr_d × epochs_d) running. + +**COMPLETE (fail-closed at M10), 21:35** — Answer to the primary question: **YES, D_MPC+ distillation visibly changes the raw policy** (SOCP-positive rate at MPC contexts up in 15/16 blocks, up to .106→.482; target-recovery RMSE down in 16/16; movement present even when the block's CFM loss rose) — **but no swept dose is globally beneficial**: 0/32 pre/post cells pass the pre-registered liveness gate on the M10 bank (r0 SR .657; post-block round-1 SR .243–.600), the base margin/e10 recipe collapses by r2–3 with or without distillation (no-distill control confirms), and per the pre-registration the M50/M100 confirmations are not triggered. Full record: MPC_STUDY_DELIVERY.json, mpc_distill/MPC_M10_SELECTION.json (all 37 cells), per-block audits in mpc_distill/M*/metrics.jsonl, every pre/post checkpoint kept. Synthesis: across deterministic v1/v2 escapes and the privileged MPC pool alike, hard-context-only distillation trades global liveness for local certifiable support; the confirmed winning recipe worked by EMBEDDING certified targets in the full replay mixture (~5% mass) — composition, not target quality, is the binding constraint. + +### Stage B extension (declared 07:55 before reading its results) + +- B0: codex margin/original/α.01/e100 checkpoints r1–r4 evaluated on the SAME M25 qual bank (matched-round control; identical recipe lineage, same commit and seeds). +- B9: margin/orig_plus_recovery/α.01/e100/lr1e-4/ess.5, rounds 1–4 from exact r0 — tests the marginal contribution of certified recovery positives ON TOP of the strongest known recipe. Comparison B0 vs B9 at matched rounds on the same bank is the pre-registered arm-2/arm-3 style contrast for the final freeze decision; freeze criterion remains the predeclared Stage-B rule. diff --git a/overnight_run_07_12_sfm/claude_recipe_study/SELECTION_RULE.json b/overnight_run_07_12_sfm/claude_recipe_study/SELECTION_RULE.json new file mode 100644 index 0000000..8a5dfa6 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/SELECTION_RULE.json @@ -0,0 +1,35 @@ +{ + "status": "PREDECLARED_BEFORE_FINAL_STUDY", + "declared_at": "2026-07-26T05:05:00-07:00", + "applies_to": "Stage D fixed-recipe per-round raw M50 evaluation ONLY", + "bank": {"ep0": 320000, "noise_seed": 20260729, "m_per_gamma": 50, "temperature": 1.0, "scene_profile": "double_density_velocity_ood"}, + "evaluator": "sfm_b1_offline_eval.py (executed sliding-window Validity, exact GREEN verifier)", + "rule": { + "candidates": "rounds r >= 1 of the single frozen recipe run", + "liveness_gate": [ + "pooled SR(r) >= pooled SR(r0) - 0.02 on the same M50 CRN bank", + "pooled timeout(r) <= pooled timeout(r0) + 0.05" + ], + "objective_order_among_eligible": [ + "1. minimize pooled CR", + "2. maximize pooled window Validity (mean of per-trajectory fractions)", + "3. maximize pooled successful minimum clearance", + "4. minimize pooled successful time-to-goal", + "5. smallest round index" + ], + "four_user_metrics_operationalization": { + "collision_rate": "objective 1 (primary)", + "validity": "objective 2", + "successful_min_clearance": "objective 3", + "time_to_goal": "objective 4; a slower checkpoint is acceptable only if it strictly wins an earlier objective, which is exactly what the lexicographic order encodes" + }, + "if_no_round_passes_gate": "report honestly that the fixed recipe did not produce an eligible improvement; publish the full Pareto frontier over (CR, Validity, clearance, time) for all rounds and do NOT promote a collapsed checkpoint", + "pareto_reporting": "alongside the winner, all non-dominated rounds on (CR down, Validity up, clearance up, time down) are listed in the delivery" + }, + "prohibitions": [ + "no per-gamma or global temperature tuning (temperature fixed 1.0)", + "no gathering-controller SR or training-loss in selection", + "no checkpoint selection from the M100 confirmation bank", + "no recipe or Kazuki changes after reading confirmation" + ] +} diff --git a/overnight_run_07_12_sfm/claude_recipe_study/STAGE_D_SELECTION.json b/overnight_run_07_12_sfm/claude_recipe_study/STAGE_D_SELECTION.json new file mode 100644 index 0000000..1dd092a --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/STAGE_D_SELECTION.json @@ -0,0 +1,32 @@ +{ + "status": "STAGE_D_SELECTION_APPLIED", + "applied_at": "2026-07-26T13:05:00-07:00", + "rule": "SELECTION_RULE.json (predeclared 2026-07-26 05:05 PT, before the final run)", + "bank": {"ep0": 320000, "noise_seed": 20260729, "m_per_gamma": 50}, + "recipe": "FIXED_RECIPE.json (cost/hard, alpha .01, exposures 10, lr 1e-4, ESS .5, rounds 10)", + "r0": {"SR": 0.720, "CR": 0.269, "timeout": 0.011, "Validity": 0.638, "clearance": 0.134, "time": 8.31}, + "per_round_pooled": { + "r1": {"SR": 0.720, "CR": 0.280, "timeout": 0.000, "Validity": 0.603, "clearance": 0.132, "time": 8.51}, + "r2": {"SR": 0.683, "CR": 0.317, "timeout": 0.000, "Validity": 0.593, "clearance": 0.139, "time": 8.06}, + "r3": {"SR": 0.683, "CR": 0.317, "timeout": 0.000, "Validity": 0.573, "clearance": 0.137, "time": 7.66}, + "r4": {"SR": 0.683, "CR": 0.317, "timeout": 0.000, "Validity": 0.571, "clearance": 0.138, "time": 7.53}, + "r5": {"SR": 0.686, "CR": 0.314, "timeout": 0.000, "Validity": 0.577, "clearance": 0.138, "time": 7.45}, + "r6": {"SR": 0.677, "CR": 0.323, "timeout": 0.000, "Validity": 0.564, "clearance": 0.144, "time": 7.29}, + "r7": {"SR": 0.737, "CR": 0.263, "timeout": 0.000, "Validity": 0.571, "clearance": 0.140, "time": 7.35}, + "r8": {"SR": 0.694, "CR": 0.306, "timeout": 0.000, "Validity": 0.558, "clearance": 0.137, "time": 7.27}, + "r9": {"SR": 0.634, "CR": 0.366, "timeout": 0.000, "Validity": 0.560, "clearance": 0.144, "time": 7.35}, + "r10": {"SR": 0.603, "CR": 0.397, "timeout": 0.000, "Validity": 0.559, "clearance": 0.143, "time": 7.44} + }, + "liveness_gate": {"SR_min": 0.700, "timeout_max": 0.061}, + "eligible_rounds": ["r1", "r7"], + "selected_round": "r7", + "selected_checkpoint": "/data3/research1/claude_sfm_best_recipe_f06e8dd/stageD/final_cost_hard/round_07.pt", + "honest_assessment": "The rule-selected r7 lowers CR by only 0.006 versus r0 while losing 0.067 Validity; the frozen recipe does NOT deliver the desired joint improvement (substantially lower CR + higher Validity + higher clearance). This is reported as a null result for iteration 1.", + "pareto_frontier_on_m50_bank": [ + {"cell": "r0", "why": "highest Validity among liveness-eligible cells (.638), CR .269"}, + {"cell": "frozen r7", "why": "lowest CR among eligible frozen-recipe rounds (.263), fastest time (7.35), but Validity .571"}, + {"cell": "companion B9-r1 (margin + orig-plus-recovery-v1, e100) — diagnostic cells, not the frozen recipe", "why": "SR .729 / CR .246 / Validity .750 / clearance .162 / time 10.51: dominates r0 on CR, Validity, clearance while passing the liveness gate; motivates the pre-registered iteration 2"}, + {"cell": "companion B9-r2", "why": "Validity .816 / clearance .151 / CR .249 but SR .643 fails the gate (timeout collapse begins)"} + ], + "paper_plot": "/data3/research1/claude_sfm_best_recipe_f06e8dd/stageD/paper_trends/b1_margin50_metric_trends.{png,pdf} via safe_flow_expansion@87063d3 contract; per-cell SEs stored in cost_hard_frozen_trends_rows.jsonl" +} diff --git a/overnight_run_07_12_sfm/claude_recipe_study/TEACHER_STUDY_DELIVERY.json b/overnight_run_07_12_sfm/claude_recipe_study/TEACHER_STUDY_DELIVERY.json new file mode 100644 index 0000000..8693678 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/TEACHER_STUDY_DELIVERY.json @@ -0,0 +1,22 @@ +{ + "status": "UNVERIFIED_MPC_TEACHER_STUDY_COMPLETE_FAIL_CLOSED_AT_ROUND_1", + "completed_at": "2026-07-28T05:30:00-07:00", + "branch": "agent/claude-sfm-unverified-teacher-run-20260728 @ pinned 2a873327 (+3 additive commits: pre-registration, authenticated default-off --reuse-buffer, canonical-content assertion; teacher tool semantics untouched)", + "headline": "The teacher never got a viable substrate: one UNCHANGED ordinary e10 replay round from r1 (theta_half, before any teacher) already collapses liveness on the dev bank (SR .77->.44, timeout .04->.43) while CR falls to .129 and Validity rises to .836; the nine pre-registered teacher doses (lr 1e-6..1e-5 x 1..4 epochs) move the policy by epsilon on top and repair nothing. 0/10 candidates admissible; loop stopped at round 1 per pre-registration; M50 (420000) and M100 (430000) banks remain untouched; no Kazuki-win claim.", + "mechanistic_finding": "conflict probes (tool-native, identical seeds/records before/after): cos(teacher gradient, ordinary D+ gradient) = 0.034 before the block, drifting to -0.07 as the teacher fits; cos(teacher, D-) = 0.377 throughout. The unverified privileged-MPC teacher's plans resemble the policy's REJECTED executed windows (fast raw plans at hard contexts) far more than its certified ones — distilling it pushes against the alpha-signed ordinary update. This is direct gradient-level evidence for why unverified MPC distillation cannot rescue this lineage.", + "d_mpc_integrity": { + "identical_across_arms": "canonical content hash cee3ad65e962473a... identical across all 9 arms (file-SHA equality is NOT a valid check: torch.save bytes differ across fresh-vs-reloaded object graphs for deep-equal payloads — verified and documented)", + "separation": "D_MPC never entered D/D+, GP, acquisition, or validity; SOCP audit-only; control bounds + privileged exact-SFM feasibility enforced by the pinned tool", + "harvest": "400 hard contexts, prefix-replay authenticated, 0 reconstruction mismatches" + }, + "round_1_table": "teacher_rounds/develop/metrics.jsonl (all 10 rows incl. no-teacher control, admissibility checks, probes; nothing hidden)", + "artifacts": { + "theta_half_and_arm_checkpoints": "teacher_rounds/develop/round_01/", + "d_mpc_buffers": "teacher_rounds/develop/round_01/arm_*/D_MPC.pt", + "branch_viz": "teacher_rounds/viz/teacher_branches_{lr1em06_ep1,lr1em05_ep4}.{png,json} (single shared buffer, two dose exemplars)", + "winner_paper_plot": "none — no admissible candidate exists to promote; rendering a winner plot would misrepresent the result", + "crashed_attempts_preserved": ["teacher_rounds/develop_crashed_buffer_sha (provenance-field byte divergence)", "teacher_rounds/develop_crashed2 (led to the canonical-hash correction; its 9 buffers verify content-identical)"] + }, + "prohibitions_respected": ["Codex branch untouched", "immutable r1 untouched", "ordinary gathering/replay/verifier/GP/raw evaluator unmodified", "no temperature tuning", "M50/M100 banks never read"], + "synthesis": "Fourth independent line of evidence that the r1 checkpoint sits on a hard liveness-safety frontier: (1) e100 continuation collapses at r2; (2) anchored low-dose continuation fails the timeout gate at every dose; (3) verified-MPC distillation visibly moves the policy but never passes liveness; (4) the ordinary e10 update alone now shows the same collapse in one round, and the unverified teacher gradient is anti-aligned with certified behavior. Any further push on this axis needs a different mechanism (e.g., longer evaluation horizon T, or liveness-constrained objectives), not a different data source." +} diff --git a/overnight_run_07_12_sfm/claude_recipe_study/TEACHER_STUDY_PREREGISTRATION.json b/overnight_run_07_12_sfm/claude_recipe_study/TEACHER_STUDY_PREREGISTRATION.json new file mode 100644 index 0000000..b58279f --- /dev/null +++ b/overnight_run_07_12_sfm/claude_recipe_study/TEACHER_STUDY_PREREGISTRATION.json @@ -0,0 +1,24 @@ +{ + "status": "PREDECLARED_BEFORE_ANY_TEACHER_ROUND", + "declared_at": "2026-07-28T00:40:00-07:00", + "branch": "agent/claude-sfm-unverified-teacher-run-20260728 pinned at 2a873327c8afd7dffa9a8e10e2bf957bf5203489 (teacher tooling used as-is; Codex branch, immutable r1, ordinary gathering/replay, verifier, GP, raw evaluator untouched)", + "start_checkpoint": {"path": "/data3/research1/claude_sfm_best_recipe_f06e8dd/stageB/B9_margin_origrec_e100/round_01.pt", "sha256_prefix": "141e4ae6592bf73f"}, + "round_structure": { + "ordinary_update": "unchanged executed D+/D- gathering (margin selector, contract seeds, continuation round k gathers expansion_scenarios(1+k), GP chain from archived B9 round-1 shard, ell 0.3259987743470518) followed by unchanged sfm_b1_offline_replay.replay with alpha .01, exposure_epochs 10, lr 1e-4, batch 128 -> theta_(n+1/2)", + "teacher_block": "claude_apply_unverified_mpc_teacher.py on theta_(n+1/2) with --audit-socp (audit-only), --max-contexts 400, seed 20260728; D_MPC kept separate from D/D+/GP/acquisition/validity; control bounds and the privileged controller's own exact-SFM feasibility enforced by the pinned tool", + "dose_arms": "teacher lr {1e-6, 3e-6, 1e-5} x epochs {1, 2, 4} = 9 arms per round, all applied to the SAME theta_(n+1/2) and the SAME round shard; identical deterministic D_MPC across arms is guaranteed by construction and ASSERTED via teacher_buffer_sha256 equality (arms 2-9 reuse arm 1's buffer through an additive default-off --reuse-buffer flag whose equivalence to fresh harvest is unit-tested)", + "no_teacher_control": "theta_(n+1/2) itself is evaluated alongside the 9 arms every round and competes in acceptance" + }, + "development_bank": {"ep0": 390000, "noise_seed": 20260740, "m_per_gamma": 10, "note": "same fixed-CRN M10 development bank and r1 baseline as the R1-continuation study (r1: SR .771, CR .186, timeout .043, V .752, clr .144); development-only, never confirmation"}, + "admissibility_gates_identical_to_R1_continuation_prereg": { + "SR": ">= .721 (r1_dev - .05)", "timeout": "<= .093 (r1_dev + .05)", + "CR": "<= .186 (hard)", "Validity": ">= .752 (hard)", "clearance": ">= .144 (hard)", + "gamma_trend": "clearance(g0.1) >= clearance(g1.0) AND time(g0.1) >= time(g1.0)" + }, + "acceptance": "lexicographic (lowest CR, highest Validity, highest clearance, shortest time) among admissible candidates INCLUDING the no-teacher control; accepted checkpoint seeds the next round; if nothing admissible the loop stops", + "rounds_max": 6, + "m50_promotion": {"bank": {"ep0": 420000, "noise_seed": 20260745, "m_per_gamma": 50}, "rule": "accepted-round candidates whose dose arm was admissible in >= 2 consecutive rounds (stability) are promoted; if the loop stops before any dose repeats, the single best admissible candidate (by the lexicographic rule) is promoted alone; r0 and r1 evaluated on the same bank for reference"}, + "final_m100": {"bank": {"ep0": 430000, "noise_seed": 20260746, "m_per_gamma": 100}, "rule": "exactly one candidate (best on M50 by the same lexicographic rule with the same liveness gates measured against r1's M50 values) vs r0, accepted r1, and locked Kazuki; paired scenario-cluster CIs; UNTOUCHED until that candidate is frozen; no Kazuki-win claim unless the CIs support it"}, + "diagnostics": "teacher-vs-D+/D- gradient cosines before/after every block (tool-native), harvest/audit counts, buffer SHAs, per-round M10 rows for all arms (nothing hidden), sfm_b1_teacher_branch_viz.py render per accepted candidate, paper_b1_margin50_trends.py for the winner", + "compute": "GPU3 only unless GPU1 is verifiably idle; SOCP audit on CPU" +} diff --git a/overnight_run_07_12_sfm/claude_stageA.py b/overnight_run_07_12_sfm/claude_stageA.py new file mode 100644 index 0000000..d4a423e --- /dev/null +++ b/overnight_run_07_12_sfm/claude_stageA.py @@ -0,0 +1,459 @@ +"""Stage A local failure diagnosis for the SFM B1 offline recipe study. + +Diagnostic-only tooling; trains nothing unless the ``update`` subcommand is +invoked, and never touches raw evaluation semantics. + +Subcommands +----------- +``mine`` — classify failure contexts of an archived ExecutedRoundShard. +``branch`` — instrumented closed-loop branch trace of declared episodes: + every one of the K=16 flow candidates is exact-verified (a + diagnostic superset of the B=4 budget), B-selection replicates + the round-1 acquisition (empty GP buffer, calibrated beta), + both execution selectors are evaluated, and the episode + advances with the B1 executed action (chosen selector, + raw-continuation at NVP) exactly as the offline collector. +``update`` — apply exactly one replay update (declared knobs, opt-in data + intervention) to the exact r0 checkpoint using an archived + round-1 shard, and save the updated checkpoint. +``compare`` — before/after tables from two ``branch`` traces. +""" +from __future__ import annotations + +import argparse +from collections import Counter +from concurrent.futures import ProcessPoolExecutor +import copy +import json +import os + +import numpy as np +import torch + +import _paths # noqa: F401 +import claude_offline_aug as AUG +import grid_policy_sfm as GPS +import sfm_b1_cost as BC +import sfm_b1_expand as BX +import sfm_b1_eval as BE +import sfm_b1_full_episode_audit as FA +import sfm_b1_offline_exec as OE +import sfm_b1_offline_store as OS +import sfm_b1_rbf as BR +import sfm_b1_store as BS +import sfm_metrics2 as SM +import sfm_protocol as SP +import sfm_scene as SS + + +def _write_json(path, payload): + OE._write_json(path, payload) + + +# ---------------------------------------------------------------- mine ---- + +def mine(args): + shard = OS.ExecutedRoundShard.load(args.shard) + pop_a, pop_b, stats = AUG.tag_populations(shard) + nvp = [w for w in shard.windows if w.get("nvp_context")] + collisions = [w for w in shard.windows if w.get("collision_after_action")] + traps = [w for w in shard.windows if w.get("trap_event")] + by_gamma = Counter( + str(shard.contexts[w["context_id"]]["gamma"]) for w in pop_b + ) + payload = dict( + shard=os.path.abspath(args.shard), + stats=stats, + NVP_contexts=len(nvp), + collision_windows=len(collisions), + trap_windows=len(traps), + popB_by_gamma=dict(by_gamma), + examples=dict( + nvp=[_ctx_key(shard, w) for w in nvp[:20]], + collision=[_ctx_key(shard, w) for w in collisions[:20]], + trap=[_ctx_key(shard, w) for w in traps[:20]], + ), + ) + _write_json(args.out, payload) + print(json.dumps({k: payload[k] for k in ( + "NVP_contexts", "collision_windows", "trap_windows")}, indent=1)) + + +def _ctx_key(shard, window): + context = shard.contexts[int(window["context_id"])] + return dict( + scenario=int(context["scenario_id"]), gamma=float(context["gamma"]), + step=int(context["step"]), + ) + + +# -------------------------------------------------------------- branch ---- + +@torch.no_grad() +def branch(args): + device = args.device + policy, _ = GPS.load_sfm_policy(args.checkpoint, device=device) + policy.eval() + phi_policy = copy.deepcopy(policy).eval() + for parameter in phi_policy.parameters(): + parameter.requires_grad_(False) + cfg = OE.OfflineConfig(alpha=0.0, exposure_epochs=1, rounds=1, smoke=True) + environment = SS.scene_profile(cfg.scene_profile) + pairs = [ + (int(s), float(g)) + for s in args.scenarios for g in args.gammas + ] + replicas = [ + BX.Replica( + scenario_id, gamma, n_ped=environment["n_ped"], + ped_speed_range=tuple(environment["ped_speed_range"]), + ) + for scenario_id, gamma in pairs + ] + # Round-1 acquisition state: empty GP buffer + calibrated beta, + # replicated exactly as the collector does at round 1. + gp = BR.RBFGP(float(args.ell), float(cfg.gp_lam)) + beta, ess = OE._calibrate_beta( + phi_policy, gp, replicas, cfg, device, round_i=1, + ) + traces = [] + outcomes = [] + with ProcessPoolExecutor(max_workers=args.workers) as executor: + for step in range(int(cfg.T)): + live, batch = BX._stack_prepared( + [r for r in replicas if r.alive], device, + ) + if not live: + break + windows, contexts, x0 = OE._keyed_windows( + policy, live, batch, K=cfg.K, round_i=1, step=step, + source="K", seed=cfg.seed, nfe=cfg.nfe, temp=cfg.temp, + ) + raw_windows, _, raw_x0 = OE._keyed_windows( + policy, live, batch, K=1, round_i=1, step=step, + source="raw_continuation", seed=cfg.seed, nfe=cfg.nfe, + temp=cfg.temp, + ) + raw_windows = raw_windows[:, 0] + windows_np = windows.detach().cpu().numpy() + raw_np = raw_windows.detach().cpu().numpy() + features = OE._features_from_x0( + phi_policy, windows, contexts, x0, cfg.phi_s, + ) + selected_by_context, sigmas = [], [] + for index, replica in enumerate(live): + generator = torch.Generator(device=features.device) + generator.manual_seed(OE._keyed_seed( + cfg.seed, 1, replica.scenario_id, + f"{replica.gamma:.8f}", step, "acquisition", + )) + selected, trace = gp.sequential_acquire( + features[index], cfg.B, beta, generator=generator, + ) + selected_by_context.append(list(map(int, selected))) + sigmas.append([float(r["chosen_sigma"]) for r in trace]) + # Diagnostic superset: verify ALL K candidates + the raw plan. + tasks = [] + for index, replica in enumerate(live): + prepared = replica.prepared + for k in range(cfg.K): + tasks.append(( + index, k, prepared["state"], windows_np[index, k], + prepared["ped_xy"], prepared["ped_vel"], + replica.gamma, + )) + tasks.append(( + index, -1, prepared["state"], raw_np[index], + prepared["ped_xy"], prepared["ped_vel"], replica.gamma, + )) + results = list(executor.map(SM.verify_in_worker, tasks)) + by_context = {} + for index, k, result in results: + by_context.setdefault(int(index), {})[int(k)] = result + + for index, replica in enumerate(live): + prepared = replica.prepared + rows = [] + for k in range(cfg.K): + result = by_context[index][k] + margin, _, _ = BC.nominal_hp_margin( + prepared["state"], windows_np[index, k][0], + prepared["ped_xy"], replica.gamma, + ) + rows.append(dict( + candidate_id=k, + y=int(result.get("y", 0)) if result.get("resolved") + else None, + resolved=bool(result.get("resolved")), + hp_margin=float(margin), + in_B=k in selected_by_context[index], + controls=windows_np[index, k], + result=result, + )) + # B1 execution semantics restricted to the B queried rows. + query_rows = [ + dict( + candidate_id=row["candidate_id"], + acquisition_step=selected_by_context[index].index( + row["candidate_id"], + ), + controls=row["controls"], + result=row["result"], + mode=None, + sigma=sigmas[index][ + selected_by_context[index].index( + row["candidate_id"], + ) + ], + ) + for row in rows if row["in_B"] and row["resolved"] + ] + chosen = {} + for selector in ("margin", "safemppi_cost"): + chosen[selector] = BC.select_admissible( + [dict(r) for r in query_rows], selector=selector, + state=prepared["state"], ped_xy=prepared["ped_xy"], + ped_vel=prepared["ped_vel"], gamma=replica.gamma, + ) + execute = chosen[args.selector] + raw_result = by_context[index][-1] + if execute is None: + controls = raw_np[index] + executed_y = ( + int(raw_result.get("y", 0)) + if raw_result.get("resolved") else None + ) + source = "raw_continuation" + else: + controls = np.asarray(execute["controls"], np.float32) + executed_y = int(execute["result"]["y"]) + source = f"verified_{args.selector}" + k_positive = sum(1 for r in rows if r["y"] == 1) + b_positive = sum( + 1 for r in rows if r["in_B"] and r["y"] == 1 + ) + b_admissible = sum( + 1 for r in rows + if r["in_B"] and r["y"] == 1 and r["hp_margin"] >= -1e-9 + ) + clearance, displacement = AUG._window_geometry( + dict( + state=prepared["state"], ped_xy=prepared["ped_xy"], + ped_vel=prepared["ped_vel"], + ), + controls, + ) + disagree = ( + chosen["margin"] is not None + and chosen["safemppi_cost"] is not None + and int(chosen["margin"]["candidate_id"]) + != int(chosen["safemppi_cost"]["candidate_id"]) + ) + traces.append(dict( + scenario=int(replica.scenario_id), + gamma=float(replica.gamma), step=int(step), + K_positive=int(k_positive), + B_positive=int(b_positive), + B_admissible=int(b_admissible), + NVP=execute is None, + K_pos_but_B_none=bool(k_positive > 0 and b_admissible == 0), + selector_disagreement=bool(disagree), + executed_source=source, + executed_y=executed_y, + executed_clearance=float(clearance), + executed_displacement=float(displacement), + sigma_selected=sigmas[index], + )) + BX._advance(replica, controls[0]) + FA._post_action_terminal(replica) + OE._finalize_alive(replicas) + for replica in replicas: + outcomes.append(dict( + scenario=int(replica.scenario_id), gamma=float(replica.gamma), + status=replica.status, steps=len(replica.controls), + min_clearance=float(replica.minimum_clearance), + )) + aggregate = dict( + contexts=len(traces), + NVP=sum(t["NVP"] for t in traces), + K_pos_but_B_none=sum(t["K_pos_but_B_none"] for t in traces), + selector_disagreement=sum(t["selector_disagreement"] for t in traces), + mean_K_positive=float(np.mean([t["K_positive"] for t in traces])), + mean_B_positive_fraction=float(np.mean([ + t["B_positive"] / cfg.B for t in traces + ])), + outcomes=Counter(o["status"] for o in outcomes), + beta=float(beta), calibrated_ess=float(ess), + ) + payload = dict( + checkpoint=os.path.abspath(args.checkpoint), + checkpoint_sha256=OS.sha256_file(args.checkpoint), + selector=args.selector, ell=float(args.ell), + scenarios=list(map(int, args.scenarios)), + gammas=list(map(float, args.gammas)), + aggregate={ + **{k: v for k, v in aggregate.items() if k != "outcomes"}, + "outcomes": dict(aggregate["outcomes"]), + }, + outcomes=outcomes, + traces=traces, + ) + torch.save(payload, args.out) + _write_json( + args.out + ".summary.json", + {k: payload[k] for k in ( + "checkpoint", "checkpoint_sha256", "selector", "aggregate", + "outcomes", + )}, + ) + print(json.dumps(payload["aggregate"], indent=1)) + + +# -------------------------------------------------------------- update ---- + +def update(args): + policy, _ = GPS.load_sfm_policy(args.checkpoint, device=args.device) + sha = OS.sha256_file(args.checkpoint) + if sha != OE.EXPECTED_CHECKPOINT_SHA256: + raise ValueError("update must start from the exact r0 checkpoint") + BS.configure_expansion_trainability(policy) + encoder_sha = BS.module_sha256(policy.enc_grid) + optimizer = torch.optim.Adam( + [p for p in policy.parameters() if p.requires_grad], lr=args.lr, + ) + shard = OS.ExecutedRoundShard.load(args.shard) + with ProcessPoolExecutor(max_workers=args.workers) as executor: + replay = AUG.replay_with_mode( + policy, optimizer, shard, mode=args.replay_mode, + alpha=args.alpha, exposure_epochs=args.exposure_epochs, + batch=128, device=args.device, seed=args.seed, + executor=executor, + ) + if BS.module_sha256(policy.enc_grid) != encoder_sha: + raise RuntimeError("visual encoder changed") + BX._save_checkpoint(policy, args.out, dict( + role="stageA_single_update", source_sha256=sha, + shard=os.path.abspath(args.shard), lr=args.lr, alpha=args.alpha, + exposure_epochs=args.exposure_epochs, replay_mode=args.replay_mode, + seed=args.seed, + )) + compact = { + k: replay.get(k) for k in ( + "positive_eligible", "negative_eligible", "optimizer_steps", + "module_relative_parameter_drift", "fixed_probe", + ) + } + compact["replay_intervention"] = { + k: v for k, v in replay.get("replay_intervention", {}).items() + if k != "recovery_audit" + } + audit = replay.get("replay_intervention", {}).get("recovery_audit") + if audit is not None: + compact["recovery_audit_counts"] = { + k: v for k, v in audit.items() if k != "rows" + } + _write_json(args.out + ".recovery_audit.json", audit) + _write_json(args.out + ".replay.json", dict( + replay={k: v for k, v in replay.items() if k != "epochs"}, + compact=compact, + )) + print(json.dumps(compact, indent=1, default=str)) + + +# ------------------------------------------------------------- compare ---- + +def compare(args): + before = torch.load(args.before, map_location="cpu", weights_only=False) + after = torch.load(args.after, map_location="cpu", weights_only=False) + rows = [] + outcomes_b = { + (o["scenario"], o["gamma"]): o for o in before["outcomes"] + } + outcomes_a = { + (o["scenario"], o["gamma"]): o for o in after["outcomes"] + } + for key in sorted(outcomes_b): + b, a = outcomes_b[key], outcomes_a.get(key) + traces_b = [ + t for t in before["traces"] + if (t["scenario"], t["gamma"]) == key + ] + traces_a = [ + t for t in after["traces"] + if (t["scenario"], t["gamma"]) == key + ] + rows.append(dict( + scenario=key[0], gamma=key[1], + status_before=b["status"], status_after=a and a["status"], + steps_before=b["steps"], steps_after=a and a["steps"], + NVP_before=sum(t["NVP"] for t in traces_b), + NVP_after=a and sum(t["NVP"] for t in traces_a), + B_pos_frac_before=float(np.mean([ + t["B_positive"] / 4 for t in traces_b + ])) if traces_b else None, + B_pos_frac_after=float(np.mean([ + t["B_positive"] / 4 for t in traces_a + ])) if traces_a else None, + )) + payload = dict( + before=dict( + checkpoint=before["checkpoint"], + aggregate=before["aggregate"], + ), + after=dict( + checkpoint=after["checkpoint"], aggregate=after["aggregate"], + ), + episodes=rows, + ) + _write_json(args.out, payload) + print(json.dumps(dict( + before=before["aggregate"], after=after["aggregate"], + ), indent=1)) + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + sub = parser.add_subparsers(dest="cmd", required=True) + + m = sub.add_parser("mine") + m.add_argument("--shard", required=True) + m.add_argument("--out", required=True) + + b = sub.add_parser("branch") + b.add_argument("--checkpoint", required=True) + b.add_argument("--scenarios", type=int, nargs="+", required=True) + b.add_argument("--gammas", type=float, nargs="+", required=True) + b.add_argument("--selector", default="margin", + choices=("margin", "safemppi_cost")) + b.add_argument("--ell", type=float, required=True, + help="round-1 lengthscale from the control run manifest") + b.add_argument("--workers", type=int, default=16) + b.add_argument("--device", default="cuda:0") + b.add_argument("--out", required=True) + + u = sub.add_parser("update") + u.add_argument("--checkpoint", required=True) + u.add_argument("--shard", required=True) + u.add_argument("--lr", type=float, default=1e-4) + u.add_argument("--alpha", type=float, default=0.01) + u.add_argument("--exposure-epochs", type=int, default=10) + u.add_argument("--replay-mode", default="original", + choices=AUG.REPLAY_MODES) + u.add_argument("--seed", type=int, default=20260724 + 1_000_003) + u.add_argument("--workers", type=int, default=16) + u.add_argument("--device", default="cuda:0") + u.add_argument("--out", required=True) + + c = sub.add_parser("compare") + c.add_argument("--before", required=True) + c.add_argument("--after", required=True) + c.add_argument("--out", required=True) + + args = parser.parse_args(argv) + dict(mine=mine, branch=branch, update=update, compare=compare)[args.cmd]( + args, + ) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/claude_teacher_rounds.py b/overnight_run_07_12_sfm/claude_teacher_rounds.py new file mode 100644 index 0000000..2d09365 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_teacher_rounds.py @@ -0,0 +1,311 @@ +"""Unverified-MPC-teacher continuation rounds (pre-registered driver). + +Per round: (1) unchanged ordinary gathering + unchanged ordinary +``sfm_b1_offline_replay.replay`` (alpha .01, exposures 10, lr 1e-4, fresh +Adam per round — the accepted lineage may switch checkpoints between rounds) +producing ``theta_(n+1/2)`` and its exact round shard; (2) the pinned +``claude_apply_unverified_mpc_teacher.py`` on ``theta_(n+1/2)`` for the nine +dose arms (lr {1e-6,3e-6,1e-5} x epochs {1,2,4}); arm 1 harvests, arms 2-9 +authenticate-and-reuse the identical deterministic ``D_MPC`` buffer, and the +driver asserts ``teacher_buffer_sha256`` equality across all nine arms; +(3) fixed-CRN raw M10 for the no-teacher control and all arms; (4) the +pre-registered gates and lexicographic rule pick the accepted checkpoint; +the loop stops when nothing is admissible. +""" +from __future__ import annotations + +import argparse +from concurrent.futures import ProcessPoolExecutor +import copy +import json +import os +import subprocess +import sys +import time + +import torch + +import _paths # noqa: F401 +import claude_continuation as CC +import grid_policy_sfm as GPS +import sfm_b1_expand as BX +import sfm_b1_offline_exec as OE +import sfm_b1_offline_replay as OR +import sfm_b1_offline_store as OS +import sfm_b1_store as BS +import sfm_protocol as SP +import sfm_scene as SS + +TEACHER_LRS = (1e-6, 3e-6, 1e-5) +TEACHER_EPOCHS = (1, 2, 4) + + +def canonical_records_sha(buffer_path): + """Content hash of a D_MPC buffer's records, independent of the + torch.save byte stream (which is not deterministic across fresh-vs- + reloaded object graphs even for deep-equal payloads).""" + import hashlib + + import numpy as np + + payload = torch.load(buffer_path, map_location="cpu", weights_only=False) + digest = hashlib.sha256() + + def feed(value): + if isinstance(value, dict): + for key in sorted(value): + digest.update(str(key).encode()) + feed(value[key]) + elif isinstance(value, (list, tuple)): + digest.update(f"#{len(value)}".encode()) + for item in value: + feed(item) + elif isinstance(value, np.ndarray): + digest.update(str(value.dtype).encode()) + digest.update(str(value.shape).encode()) + digest.update(np.ascontiguousarray(value).tobytes()) + elif isinstance(value, float): + digest.update(repr(value).encode()) + else: + digest.update(str(value).encode()) + + digest.update(payload["round_shard_sha256"].encode()) + feed(payload["records"]) + return digest.hexdigest() +ORDINARY = dict(alpha=0.01, exposure_epochs=10, lr=1e-4, batch=128) +HERE = os.path.dirname(os.path.abspath(__file__)) + + +def _gather(policy, previous_shard, round_index, opts, device, executor): + policy.eval() + phi_policy = copy.deepcopy(policy).eval() + for parameter in phi_policy.parameters(): + parameter.requires_grad_(False) + environment = SS.scene_profile(opts.scene_profile) + replicas = [ + BX.Replica(s, g, n_ped=environment["n_ped"], + ped_speed_range=tuple(environment["ped_speed_range"])) + for s in SP.expansion_scenarios(round_index, smoke=opts.smoke) + for g in SP.GAMMAS + ] + gp, _, _ = OE.gp_from_previous( + phi_policy, previous_shard, round_i=round_index, ell=CC.B9_ELL, + cap=OE.CAP, lam=opts.gp_lam, phi_s=opts.phi_s, device=device, + seed=opts.seed + round_index * 101, + ) + beta, ess = OE._calibrate_beta( + phi_policy, gp, replicas, opts, device, round_i=round_index, + ) + shard = OS.ExecutedRoundShard(round_index) + gather = OE.gather_offline_round( + policy, phi_policy, gp, beta, replicas, opts, shard, device, + executor, round_i=round_index, + ) + return shard, dict( + beta=float(beta), ess=float(ess), + outcomes={s: sum(o["status"] == s for o in gather["outcomes"]) + for s in ("success", "collision", "timeout")}, + NVP=int(gather["counts"].get("NVP_contexts", 0)), + ) + + +def _run_teacher(checkpoint, checkpoint_sha, shard_path, outdir, lr, epochs, + *, reuse_buffer, gpu, workers, audit): + env = dict(os.environ) + env.update(CUDA_DEVICE_ORDER="PCI_BUS_ID", CUDA_VISIBLE_DEVICES=str(gpu), + PYTHONPATH=HERE) + command = [ + sys.executable, + os.path.join(HERE, "claude_apply_unverified_mpc_teacher.py"), + "--checkpoint", checkpoint, + "--expected-checkpoint-sha256", checkpoint_sha, + "--round-shard", shard_path, + "--output-dir", outdir, + "--teacher-lr", f"{lr:g}", "--teacher-epochs", str(int(epochs)), + "--max-contexts", "400", "--seed", "20260728", + "--verifier-workers", str(int(workers)), "--device", "cuda:0", + ] + if audit: + command.append("--audit-socp") + if reuse_buffer: + command.extend(["--reuse-buffer", reuse_buffer]) + completed = subprocess.run( + command, env=env, cwd=HERE, capture_output=True, text=True, + ) + if completed.returncode != 0: + raise RuntimeError( + f"teacher block failed ({outdir}):\n{completed.stderr[-2000:]}" + ) + with open(os.path.join(outdir, "COMPLETE.json")) as stream: + return json.load(stream) + + +def run(args): + device = args.device + outdir = os.path.abspath(args.outdir) + if os.path.exists(outdir): + raise FileExistsError(outdir) + os.makedirs(outdir) + r1_baseline = CC.m10_pooled(args.r1_dev_metrics) + opts = CC.ContOpts() + accepted_path = os.path.abspath(args.r1_checkpoint) + previous_shard = OS.ExecutedRoundShard.load(args.previous_shard) + with ProcessPoolExecutor(max_workers=args.verifier_workers) as executor: + for round_k in range(1, int(args.rounds) + 1): + start = time.perf_counter() + round_index = 1 + round_k + round_dir = os.path.join(outdir, f"round_{round_k:02d}") + os.makedirs(round_dir) + policy, _ = GPS.load_sfm_policy(accepted_path, device=device) + shard, gather_info = _gather( + policy, previous_shard, round_index, opts, device, executor, + ) + shard_path = os.path.join( + outdir, "round_shards", f"round_{round_k:02d}.pt", + ) + shard.save(shard_path) + BS.configure_expansion_trainability(policy) + optimizer = torch.optim.Adam( + [p for p in policy.parameters() if p.requires_grad], + lr=ORDINARY["lr"], + ) + replay = OR.replay( + policy, optimizer, shard, alpha=ORDINARY["alpha"], + exposure_epochs=ORDINARY["exposure_epochs"], + batch=ORDINARY["batch"], device=device, + seed=opts.seed + round_index * 1_000_003, + ) + half_path = os.path.join(round_dir, "theta_half.pt") + BX._save_checkpoint(policy, half_path, dict( + round=round_k, phase="post_ordinary_pre_teacher", + accepted_parent=accepted_path, + )) + half_sha = OS.sha256_file(half_path) + del policy, optimizer + torch.cuda.empty_cache() + + arms = [] + reuse = None + for lr in TEACHER_LRS: + for epochs in TEACHER_EPOCHS: + name = f"lr{lr:g}_ep{epochs}".replace("-", "m") + arm_dir = os.path.join(round_dir, f"arm_{name}") + report = _run_teacher( + half_path, half_sha, shard_path, arm_dir, lr, epochs, + reuse_buffer=reuse, gpu=args.gpu_index, + workers=args.verifier_workers, + audit=(reuse is None), + ) + if reuse is None: + reuse = os.path.join(arm_dir, "D_MPC.pt") + arms.append(dict( + name=name, lr=lr, epochs=int(epochs), + checkpoint=report["post_teacher_checkpoint"], + buffer_sha=report["teacher_buffer_sha256"], + harvest_counts={ + k: v for k, v in report["harvest"].items() + if isinstance(v, (int, float, str)) + }, + update=report["update"], + probe_before=report["conflict_probe_before"], + probe_after=report["conflict_probe_after"], + )) + canonical = { + canonical_records_sha(os.path.join( + round_dir, f"arm_{arm['name']}", "D_MPC.pt", + )) + for arm in arms + } + if len(canonical) != 1: + raise RuntimeError( + f"D_MPC records diverged across arms: {canonical}" + ) + buffer_shas = {next(iter(canonical))} + specs = [(half_path, os.path.join(round_dir, "eval_theta_half"))] + specs += [ + (arm["checkpoint"], + os.path.join(round_dir, f"eval_{arm['name']}")) + for arm in arms + ] + metrics = CC.evaluate_checkpoints( + specs, cache_dir=os.path.join(outdir, "dev_cache"), + workers_each=args.eval_workers, wave=args.eval_wave, + gpu=args.gpu_index, + ) + rows = [dict( + name="no_teacher_control", lr=None, epochs=None, + checkpoint=half_path, m10=metrics[half_path], + )] + for arm in arms: + arm["m10"] = metrics[arm["checkpoint"]] + rows.append(arm) + for row in rows: + ok, checks = CC.admissible(row["m10"], r1_baseline) + row["admissible"] = ok + row["admissibility_checks"] = checks + admissible_rows = [r for r in rows if r["admissible"]] + selected = ( + min(admissible_rows, key=lambda r: CC.lex_key(r["m10"])) + if admissible_rows else None + ) + record = dict( + round=round_k, gather=gather_info, + shard=dict(D=len(shard.D), Dplus=len(shard.Dplus), + Dminus=len(shard.Dminus)), + ordinary_replay=dict(steps=replay["optimizer_steps"]), + theta_half=half_path, theta_half_sha256=half_sha, + buffer_sha256=next(iter(buffer_shas)), + r1_baseline=r1_baseline, + rows=[{k: v for k, v in row.items()} for row in rows], + n_admissible=len(admissible_rows), + selected=None if selected is None else dict( + name=selected["name"], + checkpoint=selected["checkpoint"], + m10=selected["m10"], + ), + wall_seconds=time.perf_counter() - start, + ) + with open(os.path.join(outdir, "metrics.jsonl"), "a") as stream: + stream.write(json.dumps(record, allow_nan=False) + "\n") + print(json.dumps(dict( + round=round_k, admissible=len(admissible_rows), + selected=None if selected is None else selected["name"], + selected_m10=None if selected is None else { + k: selected["m10"][k] + for k in ("SR", "CR", "Validity", "clearance", "time") + }, + cos_teacher_pos_before=arms[0]["probe_before"][ + "cos_teacher_positive"], + wall=record["wall_seconds"], + )), flush=True) + if selected is None: + print("STOP: no admissible candidate", flush=True) + break + accepted_path = selected["checkpoint"] + previous_shard = shard + OE._write_json(os.path.join(outdir, "COMPLETE.json"), dict( + status="TEACHER_ROUNDS_COMPLETE", + rounds_dir=outdir, + final_accepted=accepted_path, + final_accepted_sha256=OS.sha256_file(accepted_path), + )) + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--outdir", required=True) + parser.add_argument("--r1-checkpoint", required=True) + parser.add_argument("--previous-shard", required=True) + parser.add_argument("--r1-dev-metrics", required=True) + parser.add_argument("--rounds", type=int, default=6) + parser.add_argument("--verifier-workers", type=int, default=14) + parser.add_argument("--eval-workers", type=int, default=12) + parser.add_argument("--eval-wave", type=int, default=5) + parser.add_argument("--gpu-index", type=int, default=3) + parser.add_argument("--device", default="cuda:0") + args = parser.parse_args(argv) + run(args) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/claude_unfaithful_full.py b/overnight_run_07_12_sfm/claude_unfaithful_full.py new file mode 100644 index 0000000..4bb2751 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_unfaithful_full.py @@ -0,0 +1,52 @@ +"""UNFAITHFUL full launch: p=0.5, E=4, lr 1e-5, 10 rounds, M10 every round.""" +import argparse, json, os +from concurrent.futures import ProcessPoolExecutor +import torch +import _paths +import claude_afe_driver as AD +import claude_afe_fidelity as AF +import claude_continuation as CC +import claude_corrected_study as CS +import claude_unfaithful_pilot as UP +import grid_policy_sfm as GPS +import sfm_b1_expand as BX +import sfm_b1_offline_store as OS +import sfm_b1_store as BS + +ap = argparse.ArgumentParser(); ap.add_argument("--outdir", required=True) +ap.add_argument("--workers", type=int, default=28) +ap.add_argument("--eval-workers", type=int, default=12) +a = ap.parse_args() +out = os.path.abspath(a.outdir); os.makedirs(out, exist_ok=True) +payload = torch.load(UP.TEACHER, map_location="cpu", weights_only=False) +_, teacher_records = CS._teacher_records(payload) +b9 = OS.ExecutedRoundShard.load(AD.B9_SHARD) +cache = AF.VerifierCache() +policy, _ = GPS.load_sfm_policy(AD.R1, device="cuda:0") +BS.configure_expansion_trainability(policy) +opt = torch.optim.Adam([q for q in policy.parameters() if q.requires_grad], lr=1e-5) +prev = b9 +with ProcessPoolExecutor(max_workers=a.workers) as ex: + for k in range(1, 11): + ck = os.path.join(out, f"round_{k:02d}.pt") + sp = os.path.join(out, f"shard_{k:02d}.pt") + if os.path.isfile(ck): + policy, _ = GPS.load_sfm_policy(ck, device="cuda:0") + BS.configure_expansion_trainability(policy) + opt = torch.optim.Adam([q for q in policy.parameters() if q.requires_grad], lr=1e-5) + prev = AF.CertifiedQueryShard.load(sp); continue + policy.eval() + shard, ginfo = AD.gather_round(policy, prev, 1 + k, 0.5, "cuda:0", ex, cache, sp) + records, w = UP.mixed_weights([prev, shard], teacher_records, 0.5) + UP.replay(policy, opt, records, w, 4, 128, "cuda:0", AF.SEED + 77 * k) + BX._save_checkpoint(policy, ck, dict(arm="p05_ep4_full", round=k, UNFAITHFUL=True)) + print(json.dumps(dict(round=k, Dplus=ginfo["Dplus"], outcomes=ginfo["outcomes"])), flush=True) + prev = shard +specs = [(AD.R1, os.path.join(out, "eval_r0"))] + [ + (os.path.join(out, f"round_{k:02d}.pt"), os.path.join(out, f"eval_r{k}")) for k in range(1, 11)] +met = CC.evaluate_checkpoints(specs, cache_dir=os.path.join(out, "m10_cache"), + workers_each=a.eval_workers, wave=5, gpu=3) +with open(os.path.join(out, "FULL_M10.json"), "w") as s: + json.dump({("r1" if p == AD.R1 else os.path.basename(p)[:-3]): met[p] for p, _ in specs}, + s, indent=1, allow_nan=False, default=float) +print("FULL_M10_DONE", flush=True) diff --git a/overnight_run_07_12_sfm/claude_unfaithful_pilot.py b/overnight_run_07_12_sfm/claude_unfaithful_pilot.py new file mode 100644 index 0000000..f1dd711 --- /dev/null +++ b/overnight_run_07_12_sfm/claude_unfaithful_pilot.py @@ -0,0 +1,109 @@ +"""UNFAITHFUL pilot (explicitly directed): mix UNVERIFIED goal-directed +controller windows (no SOCP gate) into replay at portion p, corrected +whole-dataset low-dose steps, to keep goal-approach while gaining safety. +Arms: p in {0.3, 0.5} x epochs in {2, 4}, lr 1e-5, 3 rounds each from r1. +Labeled UNFAITHFUL everywhere; never mixed into certified stores. +""" +from __future__ import annotations + +import argparse +from concurrent.futures import ProcessPoolExecutor +import json +import os + +import numpy as np +import torch + +import _paths # noqa: F401 +import claude_afe_driver as AD +import claude_afe_fidelity as AF +import claude_continuation as CC +import claude_corrected_study as CS +import grid_policy_sfm as GPS +import sfm_b1_expand as BX +import sfm_b1_offline_store as OS +import sfm_b1_store as BS + +TEACHER = ("/data3/research1/claude_sfm_corrected_privileged_distill_" + "736fa6c/phase12/D_MPC_corrected.pt") + + +def mixed_weights(shards, teacher_records, p): + cert = [(s, r) for s in shards for r in s.Dplus] + hm_c, _ = BS.hierarchy_mass(cert) + hm_t, _ = BS.hierarchy_mass(teacher_records) + w = {k: (1 - p) * v for k, v in hm_c.items()} + w.update({k: p * v for k, v in hm_t.items()}) + return cert + teacher_records, w + + +def replay(policy, opt, records, w, epochs, batch, device, seed): + policy.train() + for e in range(epochs): + opt.zero_grad(set_to_none=True) + for st in range(0, len(records), batch): + v = records[st:st + batch] + g, l, h, c = BS._tensor_batch(v, device) + ww = torch.as_tensor( + [len(v) * w[(id(a), int(r["query_id"]))] for a, r in v], + dtype=c.dtype, device=device) + torch.manual_seed(seed + e * 999983 + st) + policy.cfm_loss(c, policy.ctx_from(g, l, h), + weights=ww).backward() + opt.step() + policy.eval() + + +def run(args): + out = os.path.abspath(args.outdir) + os.makedirs(out, exist_ok=True) + device = "cuda:0" + payload = torch.load(TEACHER, map_location="cpu", weights_only=False) + holder, teacher_records = CS._teacher_records(payload) + b9 = OS.ExecutedRoundShard.load(AD.B9_SHARD) + cache = AF.VerifierCache() + arms = [dict(p=p, ep=ep) for p in (0.3, 0.5) for ep in (2, 4)] + with ProcessPoolExecutor(max_workers=args.workers) as ex: + for a in arms: + name = f"p{a['p']}_ep{a['ep']}".replace(".", "") + policy, _ = GPS.load_sfm_policy(AD.R1, device=device) + BS.configure_expansion_trainability(policy) + opt = torch.optim.Adam( + [q for q in policy.parameters() if q.requires_grad], + lr=1e-5) + prev = b9 + for k in range(1, 4): + sp = os.path.join(out, f"{name}_shard{k}.pt") + shard, ginfo = AD.gather_round( + policy, prev, 1 + k, 0.5, device, ex, cache, sp) + records, w = mixed_weights( + [prev, shard], teacher_records, a["p"]) + replay(policy, opt, records, w, a["ep"], 128, device, + AF.SEED + k) + ck = os.path.join(out, f"{name}_r{k}.pt") + BX._save_checkpoint(policy, ck, dict(arm=name, round=k, + UNFAITHFUL=True)) + print(json.dumps(dict(arm=name, round=k, + Dplus=ginfo["Dplus"])), flush=True) + prev = shard + specs = [(AD.R1, os.path.join(out, "eval_r1"))] + [ + (os.path.join(out, f"p{p}_ep{e}".replace(".", "") + f"_r{k}.pt"), + os.path.join(out, f"eval_p{p}_ep{e}_r{k}".replace(".", ""))) + for p in (0.3, 0.5) for e in (2, 4) for k in (1, 3)] + met = CC.evaluate_checkpoints( + specs, cache_dir=os.path.join(out, "m10_cache"), + workers_each=args.eval_workers, wave=5, gpu=3) + with open(os.path.join(out, "PILOT.json"), "w") as s: + json.dump({os.path.basename(p): met[p] for p, _ in specs}, s, + indent=1, allow_nan=False, default=float) + print(json.dumps({os.path.basename(p): { + k: met[p][k] for k in ("SR", "CR", "Validity", "clearance", "time")} + for p, _ in specs}, allow_nan=False, default=float), flush=True) + + +if __name__ == "__main__": + ap = argparse.ArgumentParser() + ap.add_argument("--outdir", required=True) + ap.add_argument("--workers", type=int, default=28) + ap.add_argument("--eval-workers", type=int, default=12) + run(ap.parse_args()) diff --git a/overnight_run_07_12_sfm/claude_unverified_mpc_teacher.py b/overnight_run_07_12_sfm/claude_unverified_mpc_teacher.py new file mode 100644 index 0000000..e809a3f --- /dev/null +++ b/overnight_run_07_12_sfm/claude_unverified_mpc_teacher.py @@ -0,0 +1,491 @@ +"""Privileged Codex MPC teacher data kept separate from certified Safe Expansion. + +This module is an explicit ablation. It does *not* change ordinary +``D/D+/D-`` gathering, the RBF-GP, acquisition, verifier labels, or validity +accounting. At declared hard contexts it asks the historical privileged SFM +MPC controller for one bounded H=10 recovery plan and stores that plan in a +separate ``D_MPC`` buffer even when the compact SOCP would reject it. + +The intended round ordering is:: + + ordinary executed D+/D- replay + -> save theta_(n+1/2) + -> dedicated D_MPC CFM block + -> save theta_(n+1) + +The compact SOCP may be run after selection as an audit only. Its result never +gates storage and never turns a teacher record into a safety-positive sample. +The privileged controller is never used during raw evaluation. +""" +from __future__ import annotations + +from collections import defaultdict +from dataclasses import asdict +import hashlib +import json +import math +import os + +import numpy as np +import torch + +import _paths # noqa: F401 +import claude_mpc_pool as MP +import sfm_b1_offline_store as OS +import sfm_kazuki as KZ +import sfm_metrics2 as SM +import sfm_scene as SS + + +BUFFER_STATUS = "SFM_UNVERIFIED_MPC_TEACHER_BUFFER_COMPLETE" +BUFFER_VERSION = 1 +TEACHER_SOURCE = "codex_privileged_sfm_mpc" +R1_CHECKPOINT_SHA256 = ( + "141e4ae6592bf73f500ae4c5382258c8463de8519c2c7cf26f5c68ea2563b06a" +) +H = 10 +CONTROL_TOL = 1.0e-6 + + +def _controller_config_hash(): + payload = json.dumps( + asdict(MP.privileged_sfm_config()), + sort_keys=True, + separators=(",", ":"), + ).encode() + return hashlib.sha256(payload).hexdigest() + + +def _context_snapshot(context): + return { + "scenario_id": int(context["scenario_id"]), + "gamma": float(context["gamma"]), + "step": int(context["step"]), + "state": np.asarray(context["state"], np.float32).copy(), + "ped_xy": np.asarray(context["ped_xy"], np.float32).copy(), + "ped_vel": np.asarray(context["ped_vel"], np.float32).copy(), + } + + +def validate_controls(controls, *, u_max=SS.U_MAX, tol=CONTROL_TOL): + """Validate a controller command without silently repairing a bad plan.""" + raw = np.asarray(controls, np.float32) + if raw.shape != (H, 2) or not np.isfinite(raw).all(): + raise ValueError("D_MPC requires finite controls with shape [10,2]") + maximum = float(np.max(np.abs(raw))) if raw.size else 0.0 + clipped = np.clip(raw, -float(u_max), float(u_max)) + delta = float(np.max(np.abs(clipped - raw))) if raw.size else 0.0 + if delta > float(tol): + raise ValueError( + f"privileged controller exceeded U_MAX by {delta:.6g}; " + "refusing to silently reinterpret the teacher" + ) + return clipped.astype(np.float32, copy=True), { + "u_max": float(u_max), + "max_abs_before": maximum, + "max_clip_delta": delta, + } + + +def _balanced_context_ids(shard, hard_windows, max_contexts): + """Deterministic gamma-balanced subset of declared hard contexts.""" + grouped = defaultdict(list) + for window in hard_windows: + context_id = int(window["context_id"]) + context = shard.contexts[context_id] + grouped[round(float(context["gamma"]), 8)].append(context_id) + grouped = { + gamma: sorted(set(values)) + for gamma, values in grouped.items() + if values + } + if not grouped: + return [] + if max_contexts is None: + return sorted(value for values in grouped.values() for value in values) + limit = int(max_contexts) + if limit <= 0: + raise ValueError("max_contexts must be positive or None") + selected = [] + active = {gamma: list(values) for gamma, values in sorted(grouped.items())} + while active and len(selected) < limit: + next_active = {} + for gamma, values in active.items(): + if len(selected) >= limit: + break + # Evenly consume each lineage instead of taking only early steps. + index = (len(values) - 1) // 2 + selected.append(values.pop(index)) + if values: + next_active[gamma] = values + active = next_active + return sorted(set(selected)) + + +def _selected_family(pool, selected_plan): + sources = list(pool.get("candidate_sources") or ()) + for plan, source in zip(pool["plans"], sources): + if np.allclose( + np.asarray(plan, np.float32), + np.asarray(selected_plan, np.float32), + atol=1.0e-6, + ): + return dict(source) + return {"family": "controller_generated_or_escalated", "source_index": -1} + + +def select_teacher_plan(policy, context, humans, *, device, seed_step): + """Run the same local privileged SFM MPC selector used by the Codex wrapper.""" + pool = MP.build_codex_pool( + policy, + context, + humans, + device=device, + seed_step=seed_step, + track_sources=True, + ) + gamma = float(context["gamma"]) + cfg = KZ._gamma_controller_config( + MP.privileged_sfm_config(), gamma, + ).validate() + state = np.asarray(context["state"], np.float32) + ped_xy = np.asarray(context["ped_xy"], np.float32) + current_clearance = float( + np.min(np.linalg.norm(ped_xy - state[:2], axis=1) - SS.R_PED) + ) + margin = KZ._adaptive_step_filter_margin(cfg, gamma) + clearance_target = KZ._adaptive_step_filter_clearance_target( + cfg, gamma, margin, + ) + action, diagnostics, selected_plan = KZ.exact_sfm_horizon_filter_action( + humans, + state, + pool["nominal_plan"], + pool["refined_pool"], + margin=margin, + horizon=int(cfg.step_filter_horizon), + n_goal_plans=int(cfg.step_filter_goal_plans), + n_avoid_plans=int(cfg.step_filter_avoid_plans), + always_select=bool(cfg.step_filter_always_select), + min_progress=float(cfg.step_filter_min_progress), + goal_score_weight=float(cfg.step_filter_goal_score_weight), + clearance_weight=float(cfg.step_filter_clearance_weight), + fallback_clearance=bool(current_clearance < float(margin)), + fallback_lookahead=int(cfg.step_filter_fallback_lookahead), + viability_lookahead=int(cfg.step_filter_viability_lookahead), + viability_band=float(cfg.step_filter_viability_band), + viability_goal_weight=float(cfg.step_filter_viability_goal_weight), + viability_escalate=bool(cfg.step_filter_viability_escalate), + viability_escalation_band=float( + cfg.step_filter_viability_escalation_band + ), + viability_escalation_min_progress=float( + cfg.step_filter_viability_escalation_entry_progress + ), + clearance_target=clearance_target, + clearance_target_weight=float( + cfg.step_filter_clearance_target_weight + ), + ) + selected_plan, clip = validate_controls(selected_plan) + if not bool(diagnostics.get("filter_feasible")): + return None, { + "reason": "privileged_filter_infeasible", + "selector_diagnostics": diagnostics, + "pool_manifest": pool["pool_manifest"], + } + return selected_plan, { + "first_action": np.asarray(action, np.float32).tolist(), + "selector_diagnostics": diagnostics, + "candidate_source": _selected_family(pool, selected_plan), + "pool_manifest": dict(pool["pool_manifest"]), + "clip": clip, + } + + +def harvest_round( + policy, + shard, + hard_windows, + *, + device, + environment, + max_contexts=None, + executor=None, + audit_socp=False, +): + """Harvest one control-bounded privileged teacher per hard context. + + ``audit_socp`` only annotates the selected teacher. A resolved negative, + an error, and a positive are all handled identically for storage. + """ + if audit_socp and executor is None: + raise ValueError("audit_socp requires an executor") + context_ids = _balanced_context_ids(shard, hard_windows, max_contexts) + records = [] + counts = defaultdict(int) + controller_hash = _controller_config_hash() + for context_id in context_ids: + context = shard.contexts[int(context_id)] + counts["hard_contexts"] += 1 + humans, replay_state = MP.replay_prefix_humans( + shard, + int(context["scenario_id"]), + float(context["gamma"]), + int(context["step"]), + environment, + ) + ped_xy_live, _ = SS.collect_humans(humans) + if not ( + np.allclose( + replay_state, + np.asarray(context["state"], np.float32), + atol=1.0e-4, + ) + and np.allclose( + ped_xy_live, + np.asarray(context["ped_xy"], np.float32), + atol=1.0e-4, + ) + ): + counts["replay_mismatch"] += 1 + continue + plan, selection = select_teacher_plan( + policy, + context, + humans, + device=device, + seed_step=int(context["step"]), + ) + if plan is None: + counts["privileged_infeasible"] += 1 + continue + audit = None + if audit_socp: + task = ( + 0, + 0, + context["state"], + plan, + context["ped_xy"], + context["ped_vel"], + context["gamma"], + ) + _, _, result = next(iter(executor.map(SM.verify_in_worker, [task]))) + audit = { + "resolved": bool(result.get("resolved")), + "verifier_label": ( + int(result["y"]) if result.get("resolved") else None + ), + "full_h": bool(result.get("full_h", False)), + "error": result.get("error"), + "diagnostics": dict(result.get("diagnostics") or {}), + } + counts[ + "socp_audit_positive" + if audit["verifier_label"] == 1 + else "socp_audit_nonpositive" + ] += 1 + record = { + "teacher_id": len(records), + "round": int(shard.round_i), + "context_id": int(context_id), + "scenario_id": int(context["scenario_id"]), + "gamma": float(context["gamma"]), + "episode_id": int(context["scenario_id"]), + "step": int(context["step"]), + "controls": plan, + "source": TEACHER_SOURCE, + "candidate_source": selection["candidate_source"], + "controller_config_hash": controller_hash, + "selector_diagnostics": selection["selector_diagnostics"], + "pool_manifest": selection["pool_manifest"], + "clip": selection["clip"], + "context_snapshot": _context_snapshot(context), + "socp_audit": audit, + } + forbidden = {"y", "query_id", "train_eligible", "x0"} & set(record) + if forbidden: + raise AssertionError(f"D_MPC leaked ordinary-D fields: {forbidden}") + records.append(record) + counts["kept"] += 1 + per_gamma = defaultdict(int) + per_family = defaultdict(int) + for record in records: + per_gamma[f"{float(record['gamma']):g}"] += 1 + per_family[str(record["candidate_source"]["family"])] += 1 + return records, { + "counts": dict(counts), + "per_gamma": dict(per_gamma), + "per_family": dict(per_family), + "audit_socp": bool(audit_socp), + "storage_gate": ( + "privileged exact-SFM filter_feasible and bounded controls only; " + "compact SOCP is audit-only" + ), + } + + +def save_buffer(path, round_shard_path, shard, records, audit): + """Atomically save a teacher buffer authenticated to one round shard.""" + path = os.path.abspath(os.fspath(path)) + round_shard_path = os.path.abspath(os.fspath(round_shard_path)) + if int(shard.round_i) <= 0: + raise ValueError("teacher buffer requires a positive expansion round") + forbidden = {"y", "query_id", "train_eligible", "x0"} + for record in records: + leaked = forbidden & set(record) + if leaked: + raise ValueError(f"D_MPC record leaked ordinary-D fields: {leaked}") + context_id = int(record["context_id"]) + if not 0 <= context_id < len(shard.contexts): + raise ValueError("D_MPC record references a missing context") + validate_controls(record["controls"]) + payload = { + "status": BUFFER_STATUS, + "version": BUFFER_VERSION, + "round": int(shard.round_i), + "round_shard_path": round_shard_path, + "round_shard_sha256": OS.sha256_file(round_shard_path), + "records": list(records), + "audit": dict(audit), + "semantics": { + "ordinary_D_unchanged": True, + "gp_acquisition_unchanged": True, + "teacher_is_safety_label": False, + "teacher_used_at_raw_evaluation": False, + }, + } + os.makedirs(os.path.dirname(path), exist_ok=True) + temporary = path + ".tmp" + torch.save(payload, temporary) + os.replace(temporary, path) + return payload + + +def _teacher_mass(shard, records): + grouped = defaultdict(lambda: defaultdict(lambda: defaultdict(list))) + for record in records: + context = shard.contexts[int(record["context_id"])] + gamma = round(float(context["gamma"]), 8) + episode = int(context["scenario_id"]) + context_id = int(context["context_id"]) + grouped[gamma][episode][context_id].append(record) + mass = {} + if not grouped: + return mass, {"total": 0.0, "gamma": {}} + gamma_mass = defaultdict(float) + for gamma, episodes in grouped.items(): + for _, contexts in episodes.items(): + for _, values in contexts.items(): + value = ( + 1.0 + / len(grouped) + / len(episodes) + / len(contexts) + / len(values) + ) + for record in values: + mass[int(record["teacher_id"])] = value + gamma_mass[f"{gamma:g}"] += value + total = float(sum(mass.values())) + if not math.isclose(total, 1.0, rel_tol=0.0, abs_tol=1.0e-9): + raise RuntimeError(f"teacher hierarchy mass sums to {total}, not one") + return mass, {"total": total, "gamma": dict(gamma_mass)} + + +def _tensor_batch(shard, records, device): + contexts = [shard.contexts[int(record["context_id"])] for record in records] + hp10 = torch.as_tensor( + np.stack([context["hp10"] for context in contexts]), device=device, + ).float() + low = torch.as_tensor( + np.stack([context["low5"] for context in contexts]), device=device, + ).float() + hist = torch.as_tensor( + np.stack([context["hist"] for context in contexts]), device=device, + ).float() + controls = torch.as_tensor( + np.stack([record["controls"] for record in records]), device=device, + ).float() + return hp10, low, hist, controls + + +def distill_block( + policy, + optimizer, + shard, + records, + *, + epochs, + batch, + seed, +): + """Dedicated teacher-only CFM update with fresh bases every exposure.""" + if int(epochs) < 0 or int(batch) <= 0: + raise ValueError("epochs must be non-negative and batch positive") + records = list(records) + if not records or int(epochs) == 0: + return { + "steps": 0, + "optimizer_steps": 0, + "sample_exposures": 0, + "records": len(records), + "epochs": int(epochs), + "loss_first": None, + "loss_last": None, + "mass": {"total": 0.0, "gamma": {}}, + } + mass, accounting = _teacher_mass(shard, records) + rng = np.random.default_rng(int(seed)) + device = next(policy.parameters()).device + losses = [] + base_rng_states = [] + policy.train() + for epoch in range(int(epochs)): + order = list(rng.permutation(len(records))) + optimizer.zero_grad(set_to_none=True) + epoch_loss = 0.0 + for start in range(0, len(order), int(batch)): + chunk = [records[index] for index in order[start:start + int(batch)]] + hp10, low, hist, controls = _tensor_batch(shard, chunk, device) + context = policy.ctx_from(hp10, low, hist) + chunk_mass = float(sum( + mass[int(record["teacher_id"])] for record in chunk + )) + weights = torch.as_tensor( + [ + mass[int(record["teacher_id"])] + for record in chunk + ], + dtype=controls.dtype, + device=device, + ) + exposure_seed = int(seed) + epoch * 1_000_003 + start + torch.manual_seed(exposure_seed) + base_rng_states.append(exposure_seed) + loss = policy.cfm_loss(controls, context, weights=weights) + if not bool(torch.isfinite(loss)): + raise FloatingPointError("non-finite D_MPC distillation loss") + # policy.cfm_loss returns a normalized weighted mean. Multiplying + # by this chunk's global hierarchy mass and accumulating every + # chunk before one Adam step yields the declared whole-buffer + # gamma->episode->context->teacher objective exactly. + scaled = loss * chunk_mass + scaled.backward() + epoch_loss += float(scaled.detach()) + optimizer.step() + losses.append(epoch_loss) + policy.eval() + return { + "steps": len(losses), + "optimizer_steps": len(losses), + "sample_exposures": len(records) * int(epochs), + "records": len(records), + "epochs": int(epochs), + "loss_first": losses[0], + "loss_last": losses[-1], + "loss_mean": float(np.mean(losses)), + "mass": accounting, + "fresh_base_seed_count": len(set(base_rng_states)), + } diff --git a/overnight_run_07_12_sfm/recover_sfm_b1_offline_evaluation.py b/overnight_run_07_12_sfm/recover_sfm_b1_offline_evaluation.py new file mode 100644 index 0000000..b20f561 --- /dev/null +++ b/overnight_run_07_12_sfm/recover_sfm_b1_offline_evaluation.py @@ -0,0 +1,304 @@ +#!/usr/bin/env python3 +"""Recover only the raw-M50 phase of a completed offline 9-arm sweep.""" +from __future__ import annotations + +import argparse +from datetime import datetime, timezone +import json +from pathlib import Path +from types import SimpleNamespace +import time + +import run_sfm_b1_offline_9arm as RUN +import run_sfm_b1_r2_9arm as BASE + + +def _read_json(path: Path) -> dict: + if not path.is_file(): + raise RuntimeError(f"missing required artifact: {path}") + with path.open() as stream: + return json.load(stream) + + +def _validate_common_r0(path: Path) -> dict: + payload = _read_json(path) + records = payload.get("records", []) + if ( + payload.get("status") != RUN.EVAL_STATUS + or payload.get("scene_profile") != RUN.SCENE_PROFILE + or len(records) != 1 + or int(records[0].get("round", -1)) != 0 + or records[0].get("cell", {}).get("checkpoint_sha256") + != RUN.CHECKPOINT_SHA256 + ): + raise RuntimeError(f"common r0 evaluation contract mismatch: {path}") + return payload + + +def _load_frozen_training(run_root: Path, recovery_source: dict) -> tuple: + declaration_path = run_root / "RUN_DECLARATION.json" + training_path = run_root / "TRAINING_COMPLETE.json" + declaration = _read_json(declaration_path) + training_marker = _read_json(training_path) + if declaration.get("status") != "SFM_B1_OFFLINE_9ARM_DECLARED": + raise RuntimeError(f"invalid declaration status: {declaration_path}") + if training_marker.get("status") != ( + "SFM_B1_OFFLINE_9ARM_TRAINING_COMPLETE" + ): + raise RuntimeError(f"training is not complete: {training_path}") + if training_marker.get("declaration_sha256") != BASE.sha256_file( + declaration_path + ): + raise RuntimeError("training marker does not authenticate declaration") + + contract = declaration.get("contract", {}) + if declaration.get("contract_sha256") != RUN._sha256_json(contract): + raise RuntimeError("declaration contract digest mismatch") + selector = contract.get("execution_selector", "margin") + expected = { + "checkpoint_sha256": RUN.CHECKPOINT_SHA256, + "scene_profile": RUN.SCENE_PROFILE, + "rounds": RUN.ROUNDS, + "alphas": list(RUN.ALPHAS), + "exposure_epochs": list(RUN.EXPOSURE_EPOCHS), + "K": RUN.K, + "B": RUN.B, + "T": RUN.T, + "H": RUN.H, + "cap": RUN.CAP, + "gp_lambda": RUN.GP_LAMBDA, + "batch": RUN.BATCH, + "lr": RUN.LR, + "ess_target": RUN.ESS_TARGET, + "eval_M_per_gamma": 50, + "eval_temperature": 1.0, + } + for key, value in expected.items(): + if contract.get(key) != value: + raise RuntimeError( + f"frozen contract mismatch for {key}: " + f"{contract.get(key)!r} != {value!r}" + ) + if selector not in ("margin", "safemppi_cost", "balanced_rank"): + raise RuntimeError(f"unsupported frozen selector: {selector}") + if contract.get("evaluator_sha256") != BASE.sha256_file(RUN.EVALUATOR): + raise RuntimeError( + "recovery evaluator differs from the frozen evaluator" + ) + checkpoint = Path(contract["checkpoint"]).resolve() + if BASE.sha256_file(checkpoint) != RUN.CHECKPOINT_SHA256: + raise RuntimeError("frozen pretrained checkpoint digest mismatch") + + training_source = training_marker.get("source", {}) + if ( + training_source != contract.get("source") + or training_source.get("commit") is None + or recovery_source.get("commit") is None + ): + raise RuntimeError("training source provenance mismatch") + arms = list(RUN.arm_grid(selector)) + verifier_workers = int(contract["verifier_workers_per_arm"]) + seed = int(contract["seed"]) + training = { + arm.name: RUN.validate_training_arm( + run_root / "arms" / arm.name, + arm, + source_commit=training_source["commit"], + checkpoint_sha256=RUN.CHECKPOINT_SHA256, + seed=seed, + verifier_workers=verifier_workers, + ) + for arm in arms + } + return ( + declaration_path, + training_path, + contract, + training_source, + selector, + arms, + training, + checkpoint, + ) + + +def _select_recovery_gpus(args): + gpus, processes, topology = BASE.gpu_snapshot() + selected = BASE.select_idle_gpus( + gpus, + processes, + args.gpu_indices, + max_memory_mib=args.idle_memory_mib, + max_utilization=args.idle_utilization_percent, + ) + if len(selected) not in (1, 2): + raise RuntimeError( + "evaluation recovery requires one or two exclusive GPUs, got " + f"{[gpu.index for gpu in selected]}" + ) + return gpus, processes, topology, selected + + +def _recovery_allocation(arms, gpus): + if len(gpus) == 2: + return RUN.allocate_arms(arms, gpus) + if len(gpus) == 1: + return {gpus[0].uuid: list(arms)} + raise RuntimeError("evaluation recovery requires one or two GPUs") + + +def recover(args) -> dict: + started = time.perf_counter() + run_root = Path(args.run_root).resolve() + try: + run_root.relative_to(RUN.RESEARCH_ROOT.resolve()) + except ValueError as error: + raise ValueError( + f"--run-root must be below {RUN.RESEARCH_ROOT.resolve()}" + ) from error + delivery_path = run_root / "DELIVERY_COMPLETE.json" + if delivery_path.exists(): + raise FileExistsError(f"delivery already exists: {delivery_path}") + + recovery_source = BASE.source_provenance() + ( + declaration_path, + training_path, + contract, + training_source, + selector, + arms, + training, + checkpoint, + ) = _load_frozen_training(run_root, recovery_source) + runtime = SimpleNamespace( + checkpoint=str(checkpoint), + verifier_workers=int(contract["verifier_workers_per_arm"]), + seed=int(contract["seed"]), + eval_ep0=int(contract["eval_ep0"]), + eval_noise_seed=int(contract["eval_noise_seed"]), + gpu_indices=args.gpu_indices, + idle_memory_mib=int(args.idle_memory_mib), + idle_utilization_percent=int(args.idle_utilization_percent), + ) + _, _, _, gpus = _select_recovery_gpus(runtime) + allocation = _recovery_allocation(arms, gpus) + pools = BASE.allocate_cpu_pools( + arms, int(contract["verifier_workers_per_arm"]) + ) + + common_r0_dir = run_root / "evaluation" / "common_r0" + common_r0_metrics = common_r0_dir / "raw_m50_offline_metrics.json" + if common_r0_metrics.is_file(): + common_payload = _validate_common_r0(common_r0_metrics) + else: + if common_r0_dir.exists(): + raise RuntimeError( + f"refusing to overwrite partial common-r0 output: " + f"{common_r0_dir}" + ) + BASE._launch_pending( + [{ + "arm": RUN.PhaseName("common_r0"), + "gpu": gpus[0], + "cpu_pool": next(iter(pools.values())), + "command": RUN._common_r0_command(runtime, common_r0_dir), + "target": str(common_r0_dir), + }], + run_root / "logs" / "evaluation_recovery_common_r0", + ) + common_payload = _validate_common_r0(common_r0_metrics) + + jobs = RUN._phase_jobs( + runtime, arms, gpus, allocation, pools, run_root, "evaluation", + ) + pending = [] + for job in jobs: + target = Path(job["target"]) + metrics = target / "raw_m50_offline_metrics.json" + if metrics.is_file(): + continue + if target.exists(): + raise RuntimeError( + f"refusing to overwrite partial arm evaluation: {target}" + ) + pending.append(job) + if pending: + BASE._launch_pending( + pending, run_root / "logs" / "evaluation_recovery", + ) + + evaluations = { + arm.name: RUN.validate_evaluation( + run_root / "evaluation" / arm.name, + arm, + training[arm.name], + eval_ep0=runtime.eval_ep0, + eval_noise_seed=runtime.eval_noise_seed, + ) + for arm in arms + } + r0_keys = {value["r0_cell_key"] for value in evaluations.values()} + noise_hashes = { + value["noise_bank_sha256"] for value in evaluations.values() + } + if len(r0_keys) != 1 or len(noise_hashes) != 1: + raise RuntimeError("recovered evaluations do not share one raw-M50 bank") + if next(iter(r0_keys)) != ( + common_payload["records"][0]["cell"]["cell_key"] + ): + raise RuntimeError("recovered arm r0 differs from common r0") + + aggregate_dir = run_root / "evaluation" / "aggregate" + if aggregate_dir.exists(): + raise RuntimeError( + f"refusing to overwrite existing aggregate: {aggregate_dir}" + ) + aggregate_result = RUN.aggregate( + evaluations, aggregate_dir, selector=selector, + ) + manifest = { + "status": "SFM_B1_OFFLINE_9ARM_DELIVERY_COMPLETE", + "finished_at": datetime.now(timezone.utc).isoformat(), + "wall_seconds": time.perf_counter() - started, + "source": training_source, + "recovery_source": recovery_source, + "recovery_role": ( + "evaluation-only recovery; authenticated training checkpoints " + "were not modified or regenerated" + ), + "contract": contract, + "declaration": str(declaration_path), + "declaration_sha256": BASE.sha256_file(declaration_path), + "training_marker": str(training_path), + "training_marker_sha256": BASE.sha256_file(training_path), + "training": training, + "evaluations": evaluations, + "common_r0_metrics": str(common_r0_metrics), + "common_r0_metrics_sha256": BASE.sha256_file(common_r0_metrics), + "common_r0_cell_key": next(iter(r0_keys)), + "common_noise_bank_sha256": next(iter(noise_hashes)), + "aggregate": aggregate_result, + } + RUN._write_json(delivery_path, manifest) + print(json.dumps({ + "status": manifest["status"], + "selector": selector, + "wall_seconds": manifest["wall_seconds"], + "best_screening_cell": aggregate_result["best_screening_cell"], + "delivery": str(delivery_path), + }, indent=2, allow_nan=False)) + return manifest + + +def _parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--run-root", required=True) + parser.add_argument("--gpu-indices", default="1,3") + parser.add_argument("--idle-memory-mib", type=int, default=1024) + parser.add_argument("--idle-utilization-percent", type=int, default=5) + return parser + + +if __name__ == "__main__": + recover(_parser().parse_args()) diff --git a/overnight_run_07_12_sfm/run_sfm_b1_neutral_round50.sh b/overnight_run_07_12_sfm/run_sfm_b1_neutral_round50.sh new file mode 100755 index 0000000..fe01d90 --- /dev/null +++ b/overnight_run_07_12_sfm/run_sfm_b1_neutral_round50.sh @@ -0,0 +1,66 @@ +#!/usr/bin/env bash +set -euo pipefail + +if [[ $# -ne 2 ]]; then + echo "usage: $0 CHECKPOINT OUTPUT_ROOT" >&2 + exit 2 +fi + +CHECKPOINT=$(realpath "$1") +OUTPUT_ROOT=$(realpath -m "$2") +PYTHON=${PYTHON:-/home/dohyun/miniforge3/envs/cfm_mppi/bin/python} +WORKERS=${VERIFIER_WORKERS_PER_ARM:-48} +EVAL_ROUNDS=0,1,2,5,10,20,30,40,50 + +if [[ -e "$OUTPUT_ROOT" ]]; then + echo "refusing to reuse output root: $OUTPUT_ROOT" >&2 + exit 1 +fi +mkdir -p "$OUTPUT_ROOT/logs" + +names=(lr1em5_s01 lr1em5_s04 lr3em5_s01 lr3em5_s04) +lrs=(1e-5 1e-5 3e-5 3e-5) +steps=(1 4 1 4) +gpus=(1 1 3 3) +pids=() + +for index in "${!names[@]}"; do + name=${names[$index]} + ( + export CUDA_DEVICE_ORDER=PCI_BUS_ID + export CUDA_VISIBLE_DEVICES=${gpus[$index]} + export OMP_NUM_THREADS=1 + export MKL_NUM_THREADS=1 + exec "$PYTHON" overnight_run_07_12_sfm/sfm_b1_neutral_multiround.py \ + --checkpoint "$CHECKPOINT" \ + --output-root "$OUTPUT_ROOT/$name" \ + --name "$name" \ + --rounds 50 \ + --scenario-ep0 260000 \ + --eval-ep0 270000 \ + --eval-M 20 \ + --eval-rounds "$EVAL_ROUNDS" \ + --lr "${lrs[$index]}" \ + --inner-steps "${steps[$index]}" \ + --probe-per-gamma 4 \ + --device cuda \ + --workers "$WORKERS" + ) >"$OUTPUT_ROOT/logs/$name.log" 2>&1 & + pids+=("$!") +done + +failed=0 +for index in "${!pids[@]}"; do + if ! wait "${pids[$index]}"; then + echo "arm failed: ${names[$index]}" >&2 + failed=1 + fi +done +if [[ $failed -ne 0 ]]; then + exit 1 +fi + +for name in "${names[@]}"; do + test -f "$OUTPUT_ROOT/$name/DELIVERY_COMPLETE.json" +done +echo "SFM_B1_NEUTRAL_ROUND50_FOUR_ARM_COMPLETE $OUTPUT_ROOT" diff --git a/overnight_run_07_12_sfm/run_sfm_b1_offline_9arm.py b/overnight_run_07_12_sfm/run_sfm_b1_offline_9arm.py new file mode 100644 index 0000000..ba53bb4 --- /dev/null +++ b/overnight_run_07_12_sfm/run_sfm_b1_offline_9arm.py @@ -0,0 +1,1013 @@ +#!/usr/bin/env python3 +"""Frozen end-to-end launcher for the offline executed-window 9-arm study. + +The two phases are deliberately separate: + +1. train alpha {0,.01,.1} x exposure epochs {1,10,100} for ten rounds; +2. evaluate every r0--r10 checkpoint with the same raw temperature-one + M=50/gamma bank and terminal-truncated executed-window Validity. + +All nine jobs in a phase start concurrently on two exclusive GPUs with a +deterministic 5/4 allocation. Any child failure stops its peers. The +output root must not exist, so a partial study can never be mistaken for a +resumed or complete scientific run. +""" +from __future__ import annotations + +import argparse +import csv +from dataclasses import asdict, dataclass +from datetime import datetime, timezone +import json +import math +import os +from pathlib import Path +import sys +import time + +import run_sfm_b1_r2_9arm as BASE + + +HERE = Path(__file__).resolve().parent +ROOT = HERE.parent +TRAINER = HERE / "sfm_b1_offline_exec.py" +EVALUATOR = HERE / "sfm_b1_offline_eval.py" +ALPHAS = (0.0, 0.01, 0.1) +EXPOSURE_EPOCHS = (1, 10, 100) +ROUNDS = 10 +ARM_STATUS = "SFM_B1_OFFLINE_EXEC_COMPLETE" +EVAL_STATUS = "SFM_B1_OFFLINE_RAW_M50_COMPLETE" +CHECKPOINT_SHA256 = ( + "1b5179c935d3eeff8824967d707d64cc9bab273949ee1f0e4f190172bab1b215" +) +SCENE_PROFILE = "double_density_velocity_ood" +ELL_MULTIPLIER = 0.5 +CAP = 512 +GP_LAMBDA = 1.0e-2 +K = 16 +B = 4 +T = 180 +H = 10 +BATCH = 128 +LR = 1.0e-4 +ESS_TARGET = 0.5 +RESEARCH_ROOT = Path("/data3/research1") + + +def _utc_now() -> str: + return datetime.now(timezone.utc).isoformat() + + +def _write_json(path: str | os.PathLike[str], payload) -> None: + path = Path(path) + path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_name(path.name + ".tmp") + with temporary.open("w") as stream: + json.dump(payload, stream, indent=2, allow_nan=False) + os.replace(temporary, path) + + +def _sha256_json(payload) -> str: + encoded = json.dumps( + payload, sort_keys=True, separators=(",", ":"), allow_nan=False, + ).encode() + import hashlib + + return hashlib.sha256(encoded).hexdigest() + + +@dataclass(frozen=True) +class Arm: + alpha: float + exposure_epochs: int + selector: str = "margin" + + @property + def name(self) -> str: + alpha = str(float(self.alpha)).replace(".", "p") + prefixes = { + "margin": "offline_exec", + "safemppi_cost": "offline_exec_safemppi_cost", + "balanced_rank": "offline_exec_balanced_rank", + } + prefix = prefixes[self.selector] + return ( + f"{prefix}_alpha{alpha}_" + f"exposures{int(self.exposure_epochs):03d}" + ) + + +@dataclass(frozen=True) +class PhaseName: + name: str + + +def arm_grid(selector="margin") -> tuple[Arm, ...]: + if selector not in ("margin", "safemppi_cost", "balanced_rank"): + raise ValueError(f"unknown execution selector: {selector}") + return tuple( + Arm(alpha, epochs, selector) + for alpha in ALPHAS + for epochs in EXPOSURE_EPOCHS + ) + + +def _validated_output_root(value: str | os.PathLike[str]) -> Path: + path = Path(value).resolve() + research_root = RESEARCH_ROOT.resolve() + try: + path.relative_to(research_root) + except ValueError as error: + raise ValueError( + f"--outdir must be below {research_root}, got {path}" + ) from error + if path.exists(): + raise FileExistsError( + f"scientific output root must not already exist: {path}" + ) + return path + + +def allocate_arms( + arms: list[Arm], gpus: list[BASE.GPU], +) -> dict[str, list[Arm]]: + """Use both requested GPUs with a deterministic 5/4 workload split.""" + if len(gpus) != 2: + raise RuntimeError(f"exactly two idle GPUs are required, got {len(gpus)}") + selectors = {arm.selector for arm in arms} + if len(selectors) != 1 or set(arms) != set(arm_grid(next(iter(selectors)))): + raise ValueError("offline launcher requires the complete declared arm grid") + ordered_gpus = sorted(gpus, key=lambda gpu: int(gpu.index)) + allocation = {gpu.uuid: [] for gpu in ordered_gpus} + ordered_arms = sorted( + arms, key=lambda arm: (-arm.exposure_epochs, arm.alpha), + ) + for index, arm in enumerate(ordered_arms): + allocation[ordered_gpus[index % 2].uuid].append(arm) + counts = sorted(len(values) for values in allocation.values()) + if counts != [4, 5] or any(not values for values in allocation.values()): + raise RuntimeError(f"invalid two-GPU allocation: {counts}") + return allocation + + +def _trainer_command(args, arm: Arm, output: Path) -> list[str]: + return [ + sys.executable, + str(TRAINER), + "--checkpoint", + str(Path(args.checkpoint).resolve()), + "--outdir", + str(output.resolve()), + "--alpha", + str(arm.alpha), + "--exposure-epochs", + str(arm.exposure_epochs), + "--selector", + arm.selector, + "--rounds", + str(ROUNDS), + "--verifier-workers", + str(args.verifier_workers), + "--seed", + str(args.seed), + "--device", + "cuda:0", + ] + + +def _evaluation_command( + args, arm: Arm, arm_dir: Path, output: Path, *, cache_dir: Path, +) -> list[str]: + # Use the one promoted source file for r0. Per-arm round_00 containers + # embed arm-specific recipe metadata and therefore have different file + # hashes despite identical model tensors. + checkpoints = [str(Path(args.checkpoint).resolve())] + [ + str((arm_dir / f"round_{round_i:02d}.pt").resolve()) + for round_i in range(1, ROUNDS + 1) + ] + labels = [f"r{round_i}" for round_i in range(ROUNDS + 1)] + return [ + sys.executable, + str(EVALUATOR), + "--checkpoints", + *checkpoints, + "--labels", + *labels, + "--scene-profile", + SCENE_PROFILE, + "--ep0", + str(args.eval_ep0), + "--noise-seed", + str(args.eval_noise_seed), + "--device", + "cuda:0", + "--workers", + str(args.verifier_workers), + "--cache-dir", + str(cache_dir.resolve()), + "--output-dir", + str(output.resolve()), + ] + + +def _common_r0_command(args, output: Path) -> list[str]: + return [ + sys.executable, + str(EVALUATOR), + "--checkpoints", + str(Path(args.checkpoint).resolve()), + "--labels", + "r0", + "--scene-profile", + SCENE_PROFILE, + "--ep0", + str(args.eval_ep0), + "--noise-seed", + str(args.eval_noise_seed), + "--device", + "cuda:0", + "--workers", + str(args.verifier_workers), + "--cache-dir", + str((output / "cache").resolve()), + "--output-dir", + str(output.resolve()), + ] + + +def _validate_sidecar(path: Path) -> dict: + if not path.is_file(): + raise RuntimeError(f"missing artifact: {path}") + digest = BASE.sha256_file(path) + sidecar = Path(str(path) + ".COMPLETE.json") + if not sidecar.is_file(): + raise RuntimeError(f"missing artifact sidecar: {sidecar}") + with sidecar.open() as stream: + payload = json.load(stream) + if payload.get("sha256") != digest: + raise RuntimeError(f"artifact sidecar digest mismatch: {sidecar}") + return { + "path": str(path.resolve()), + "sha256": digest, + "sidecar": str(sidecar.resolve()), + "sidecar_sha256": BASE.sha256_file(sidecar), + "sidecar_payload": payload, + } + + +def _normalized_training_recipe(recipe, selector: str): + if ( + selector == "margin" + and isinstance(recipe, dict) + and "selector" not in recipe + ): + return {**recipe, "selector": "margin"} + return recipe + + +def validate_training_arm( + arm_dir: Path, + arm: Arm, + *, + source_commit: str, + checkpoint_sha256: str, + seed: int, + verifier_workers: int, +) -> dict: + marker = arm_dir / "COMPLETE.json" + if not marker.is_file(): + raise RuntimeError(f"missing arm completion marker: {marker}") + with marker.open() as stream: + payload = json.load(stream) + if payload.get("status") != ARM_STATUS: + raise RuntimeError(f"invalid arm status: {marker}") + if payload.get("experiment") != arm.name: + raise RuntimeError(f"arm identity mismatch: {marker}") + expected_recipe = { + "alpha": float(arm.alpha), + "exposure_epochs": int(arm.exposure_epochs), + "selector": arm.selector, + "rounds": ROUNDS, + "K": K, + "B": B, + "T": T, + "H": H, + "batch": BATCH, + "lr": LR, + "ess_target": ESS_TARGET, + "nfe": 8, + "temp": 1.0, + "phi_s": 0.9, + "gp_lam": GP_LAMBDA, + "verifier_workers": int(verifier_workers), + "seed": int(seed), + "scene_profile": SCENE_PROFILE, + "smoke": False, + } + observed_recipe = _normalized_training_recipe( + payload.get("recipe"), arm.selector, + ) + if observed_recipe != expected_recipe: + raise RuntimeError(f"training recipe mismatch: {marker}") + constants = payload.get("constants", {}) + ell0 = float(constants.get("ell0", -1.0)) + ell = float(constants.get("ell", -1.0)) + if ( + not math.isfinite(ell0) + or ell0 <= 0.0 + or not math.isclose( + ell, ell0 * ELL_MULTIPLIER, rel_tol=1.0e-12, abs_tol=1.0e-12, + ) + or constants.get("ell_preflight", {}).get("count") != 50 + or constants.get("ell_preflight", {}).get("representation") + != "stored proposal x0 at s=0.9" + ): + raise RuntimeError(f"invalid x0-aware ell preflight: {marker}") + expected_constants = { + "ell": ell, + "ell0": ell0, + "ell_preflight": constants["ell_preflight"], + "gp_buffer_cap": CAP, + "gp_lambda": GP_LAMBDA, + "expected_checkpoint_sha256": CHECKPOINT_SHA256, + "replay_window_rounds": 1, + "gp_quota_semantics": ( + "exactly 73 executed D+ rows per gamma plus one rotating " + "extra; any support shortage aborts the scientific round" + ), + "ess_target_semantics": ( + "mean normalized ESS over each sequential remaining pool" + ), + } + if payload.get("constants") != expected_constants: + raise RuntimeError(f"training constants mismatch: {marker}") + if payload.get("source_checkpoint_sha256") != checkpoint_sha256: + raise RuntimeError(f"source checkpoint mismatch: {marker}") + source = payload.get("source", {}) + if ( + source.get("commit") != source_commit + or source.get("tracked_worktree_clean") is not True + ): + raise RuntimeError(f"trainer source provenance mismatch: {marker}") + if payload.get("scientific_role") != ( + "offline_expansion_data_collector_not_safe_controller" + ): + raise RuntimeError(f"collector role mismatch: {marker}") + + checkpoints = [] + for round_i in range(ROUNDS + 1): + checkpoint = _validate_sidecar( + arm_dir / f"round_{round_i:02d}.pt" + ) + if checkpoint["sidecar_payload"].get("status") != "COMPLETE": + raise RuntimeError(f"invalid checkpoint sidecar: {checkpoint['sidecar']}") + checkpoints.append({"round": round_i, **checkpoint}) + + history = payload.get("history") + if not isinstance(history, list) or [ + int(row.get("round", -1)) for row in history + ] != list(range(1, ROUNDS + 1)): + raise RuntimeError(f"arm must contain rounds 1--{ROUNDS}: {marker}") + rounds = [] + for row in history: + round_i = int(row["round"]) + if row.get("experiment") != arm.name: + raise RuntimeError(f"round experiment mismatch: {marker}") + if row.get("checkpoint_sha256") != checkpoints[round_i]["sha256"]: + raise RuntimeError(f"round checkpoint digest mismatch: {marker}") + gp_selection = row.get("gp_selection", {}) + expected_gp_count = 0 if round_i == 1 else CAP + per_gamma_gp = gp_selection.get("per_gamma", {}) + if ( + int(gp_selection.get("requested_cap", -1)) != CAP + or int(gp_selection.get("quota", -1)) != CAP // 7 + or int(gp_selection.get("selected", -1)) != expected_gp_count + or sum(int(value) for value in per_gamma_gp.values()) + != expected_gp_count + or len(row.get("gp_buffer_ids", [])) != expected_gp_count + or len({ + tuple(identity) for identity in row.get("gp_buffer_ids", []) + }) != expected_gp_count + ): + raise RuntimeError(f"previous-round GP contract mismatch in round {round_i}") + if round_i > 1 and gp_selection.get("unique") is not True: + raise RuntimeError(f"GP buffer is not unique in round {round_i}") + if round_i > 1: + extra_gamma = ( + 0.1, 0.2, 0.3, 0.4, 0.5, 0.7, 1.0 + )[(round_i - 2) % 7] + expected_per_gamma = { + str(gamma): 73 + int(gamma == extra_gamma) + for gamma in (0.1, 0.2, 0.3, 0.4, 0.5, 0.7, 1.0) + } + if per_gamma_gp != expected_per_gamma: + raise RuntimeError( + f"strict gamma GP quota mismatch in round {round_i}" + ) + if "outcomes" in row: + raise RuntimeError("outcomes must be stored only inside gather") + shard = row.get("shard", {}) + shard_path = Path(shard.get("path", "")) + shard_artifact = _validate_sidecar(shard_path) + if shard_artifact["sidecar_payload"].get("status") != ( + "OFFLINE_EXECUTED_ROUND_SHARD_COMPLETE" + ): + raise RuntimeError(f"invalid executed shard sidecar: {shard_path}") + if shard_artifact["sha256"] != shard.get("sha256"): + raise RuntimeError(f"executed shard digest mismatch: {shard_path}") + gather = row.get("gather", {}) + if len(gather.get("outcomes", [])) != 56: + raise RuntimeError(f"round {round_i} must contain 56 episode outcomes") + counts = gather.get("counts", {}) + summary = gather.get("shard", {}) + contexts = int(counts.get("contexts", -1)) + if int(counts.get("B_queries", -1)) != contexts * B: + raise RuntimeError(f"B query accounting mismatch in round {round_i}") + if ( + int(summary.get("contexts", -1)) != contexts + or int(summary.get("D", -1)) != contexts + or int(summary.get("Dplus", -1)) + + int(summary.get("Dminus", -1)) != contexts + or int(summary.get("errors", -1)) != 0 + or int(summary.get("unresolved_contexts", -1)) != 0 + ): + raise RuntimeError(f"executed D partition mismatch in round {round_i}") + replay = row.get("replay", {}) + dplus = int(summary["Dplus"]) + dminus = int(summary["Dminus"]) + expected_batches = math.ceil((dplus + dminus) / BATCH) + expected_steps = expected_batches * int(arm.exposure_epochs) + if ( + replay.get("exact_once_per_exposure_epoch") is not True + or int(replay.get("positive_eligible", -1)) != dplus + or int(replay.get("negative_eligible", -1)) != dminus + or int(replay.get("positive_total_visits", -1)) + != dplus * int(arm.exposure_epochs) + or int(replay.get("negative_total_visits", -1)) + != dminus * int(arm.exposure_epochs) + or int(replay.get("optimizer_steps", -1)) != expected_steps + or bool(replay.get("negative_used_for_training")) + != bool(float(arm.alpha) > 0.0 and dminus) + ): + raise RuntimeError(f"offline replay accounting mismatch in round {round_i}") + if replay.get("visual_encoder_sha_before") != replay.get( + "visual_encoder_sha_after" + ): + raise RuntimeError(f"visual encoder changed in round {round_i}") + rounds.append({ + "round": round_i, + "D": contexts, + "Dplus": dplus, + "Dminus": dminus, + "optimizer_steps": expected_steps, + "shard": shard_artifact, + }) + return { + "arm": arm.name, + "alpha": arm.alpha, + "exposure_epochs": arm.exposure_epochs, + "marker": str(marker.resolve()), + "marker_sha256": BASE.sha256_file(marker), + "checkpoints": checkpoints, + "rounds": rounds, + } + + +def validate_evaluation( + output: Path, + arm: Arm, + training: dict, + *, + eval_ep0: int, + eval_noise_seed: int, +) -> dict: + metrics = output / "raw_m50_offline_metrics.json" + if not metrics.is_file(): + raise RuntimeError(f"missing evaluation metrics: {metrics}") + with metrics.open() as stream: + payload = json.load(stream) + if ( + payload.get("status") != EVAL_STATUS + or payload.get("scene_profile") != SCENE_PROFILE + or int(payload.get("bank", {}).get("ep0", -1)) != int(eval_ep0) + or int(payload.get("bank", {}).get("M_per_gamma", -1)) != 50 + or int(payload.get("noise_bank", {}).get("seed", -1)) + != int(eval_noise_seed) + or float(payload.get("noise_bank", {}).get("temperature", -1)) + != 1.0 + ): + raise RuntimeError(f"evaluation contract mismatch: {metrics}") + records = payload.get("records") + if not isinstance(records, list) or [ + int(row.get("round", -1)) for row in records + ] != list(range(ROUNDS + 1)): + raise RuntimeError(f"evaluation must contain r0--r{ROUNDS}: {metrics}") + expected_hashes = [CHECKPOINT_SHA256] + [ + row["sha256"] for row in training["checkpoints"][1:] + ] + for record, expected_hash in zip(records, expected_hashes): + cell = record.get("cell", {}) + if ( + cell.get("status") != "SFM_B1_OFFLINE_RAW_CELL_COMPLETE" + or cell.get("checkpoint_sha256") != expected_hash + or int(cell.get("M_per_gamma", -1)) != 50 + or int(cell.get("summary", {}).get("pooled", {}).get( + "verifier_errors", -1 + )) != 0 + ): + raise RuntimeError(f"invalid evaluation cell: {metrics}") + pooled = cell["summary"]["pooled"] + if not math.isclose( + float(pooled["SR"]) + float(pooled["CR"]) + + float(pooled["timeout"]), + 1.0, + rel_tol=0.0, + abs_tol=1.0e-12, + ): + raise RuntimeError(f"evaluation outcomes do not partition: {metrics}") + expected_outputs = [ + output / "raw_m50_offline_curves.png", + output / "raw_m50_offline_curves.pdf", + output / "raw_m50_offline_curves.figure.json", + ] + artifacts = [ + {"path": str(path.resolve()), "sha256": BASE.sha256_file(path)} + for path in [metrics, *expected_outputs] + if path.is_file() + ] + if len(artifacts) != 4: + raise RuntimeError(f"missing evaluation presentation artifact: {output}") + return { + "arm": arm.name, + "metrics": str(metrics.resolve()), + "metrics_sha256": BASE.sha256_file(metrics), + "records": records, + "noise_bank_sha256": payload["noise_bank"]["sha256"], + "r0_cell_key": records[0]["cell"]["cell_key"], + "artifacts": artifacts, + } + + +def _cell_row(arm: Arm, record: dict) -> dict: + pooled = record["cell"]["summary"]["pooled"] + clearance = pooled["successful_clearance"]["mean"] + time_to_goal = pooled["successful_time_to_goal"]["mean"] + return { + "arm": arm.name, + "alpha": float(arm.alpha), + "exposure_epochs": int(arm.exposure_epochs), + "round": int(record["round"]), + "SR": float(pooled["SR"]), + "CR": float(pooled["CR"]), + "timeout": float(pooled["timeout"]), + "Validity": float(pooled["Validity"]["mean"]), + "clearance": None if clearance is None else float(clearance), + "time_to_goal": None if time_to_goal is None else float(time_to_goal), + } + + +def _screening_key(row: dict) -> tuple: + clearance = ( + -float(row["clearance"]) + if row["clearance"] is not None else float("inf") + ) + time_to_goal = ( + float(row["time_to_goal"]) + if row["time_to_goal"] is not None else float("inf") + ) + return ( + float(row["CR"]), + -float(row["Validity"]), + -float(row["SR"]), + clearance, + time_to_goal, + int(row["round"]), + int(row["exposure_epochs"]), + float(row["alpha"]), + ) + + +def _render_aggregate( + rows: list[dict], output: Path, *, selector: str, +) -> list[dict]: + import matplotlib + + matplotlib.use("Agg") + import matplotlib.pyplot as plt + + colors = {1: "#0072B2", 10: "#E69F00", 100: "#CC79A7"} + linestyles = {0.0: "-", 0.01: "--", 0.1: ":"} + specs = ( + ("CR", "Collision rate", (-0.03, 1.03)), + ("Validity", "Validity", (-0.03, 1.03)), + ("clearance", "Min. clearance [m]", None), + ("time_to_goal", "Time-to-goal [s]", None), + ) + figure, axes = plt.subplots(2, 2, figsize=(14.5, 10.0), squeeze=False) + for axis, (key, title, ylim) in zip(axes.flat, specs): + for arm in arm_grid(selector): + values = [ + row for row in rows if row["arm"] == arm.name + ] + values.sort(key=lambda row: int(row["round"])) + axis.plot( + [row["round"] for row in values], + [ + float("nan") if row[key] is None else row[key] + for row in values + ], + color=colors[arm.exposure_epochs], + linestyle=linestyles[arm.alpha], + linewidth=1.8, + alpha=0.85, + ) + axis.set_title(title) + axis.set_xlabel("Expansion round") + axis.set_xticks(range(ROUNDS + 1)) + axis.grid(alpha=0.25) + if ylim is not None: + axis.set_ylim(*ylim) + handles = [ + plt.Line2D( + [0], [0], color=colors[epochs], lw=2.5, + label=f"{epochs} exposure epochs", + ) + for epochs in EXPOSURE_EPOCHS + ] + handles.extend( + plt.Line2D( + [0], [0], color="black", linestyle=linestyles[alpha], + lw=2.0, label=rf"$\alpha={alpha:g}$", + ) + for alpha in ALPHAS + ) + figure.legend( + handles=handles, ncol=6, loc="upper center", frameon=False + ) + figure.tight_layout(rect=(0, 0, 1, 0.93)) + artifacts = [] + for suffix in ("png", "pdf"): + path = output / f"factorial_raw_m50_pooled.{suffix}" + figure.savefig(path, dpi=300, bbox_inches="tight") + artifacts.append({ + "path": str(path.resolve()), + "sha256": BASE.sha256_file(path), + }) + plt.close(figure) + return artifacts + + +def aggregate( + evaluations: dict[str, dict], output: Path, *, selector: str, +) -> dict: + output.mkdir(parents=True, exist_ok=False) + rows = [] + for arm in arm_grid(selector): + rows.extend( + _cell_row(arm, record) + for record in evaluations[arm.name]["records"] + ) + csv_path = output / "factorial_raw_m50_metrics.csv" + fields = ( + "arm", "alpha", "exposure_epochs", "round", + "SR", "CR", "timeout", "Validity", "clearance", "time_to_goal", + ) + with csv_path.open("w", newline="") as stream: + writer = csv.DictWriter(stream, fieldnames=fields) + writer.writeheader() + writer.writerows(rows) + candidates = [row for row in rows if int(row["round"]) > 0] + best = min(candidates, key=_screening_key) + figures = _render_aggregate(rows, output, selector=selector) + result = { + "status": "SFM_B1_OFFLINE_9ARM_AGGREGATE_COMPLETE", + "selection_role": ( + "exploratory common-bank M50 screening only; not an independent " + "confirmation or a probabilistic safety guarantee" + ), + "selection_rule": ( + "post-expansion only: lower CR, higher window Validity, higher SR, " + "higher successful-only clearance, lower successful-only time, " + "then earlier round/lower exposure/lower alpha" + ), + "best_screening_cell": best, + "rows": rows, + "artifacts": [ + { + "path": str(csv_path.resolve()), + "sha256": BASE.sha256_file(csv_path), + }, + *figures, + ], + } + path = output / "AGGREGATE_COMPLETE.json" + _write_json(path, result) + result["marker"] = str(path.resolve()) + result["marker_sha256"] = BASE.sha256_file(path) + return result + + +def _select_exactly_two_gpus(args): + gpus, processes, topology = BASE.gpu_snapshot() + selected = BASE.select_idle_gpus( + gpus, + processes, + args.gpu_indices, + max_memory_mib=args.idle_memory_mib, + max_utilization=args.idle_utilization_percent, + ) + if len(selected) != 2: + raise RuntimeError( + f"the declared study requires two exclusive GPUs, got " + f"{[gpu.index for gpu in selected]}" + ) + return gpus, processes, topology, selected + + +def _phase_jobs(args, arms, selected, allocation, pools, outdir, phase): + by_uuid = {gpu.uuid: gpu for gpu in selected} + arm_gpu = { + arm: by_uuid[uuid] + for uuid, values in allocation.items() + for arm in values + } + jobs = [] + for arm in arms: + if phase == "training": + target = outdir / "arms" / arm.name + command = _trainer_command(args, arm, target) + elif phase == "evaluation": + target = outdir / "evaluation" / arm.name + command = _evaluation_command( + args, + arm, + outdir / "arms" / arm.name, + target, + cache_dir=outdir / "evaluation" / "common_r0" / "cache", + ) + else: + raise ValueError(phase) + jobs.append({ + "arm": arm, + "gpu": arm_gpu[arm], + "cpu_pool": pools[arm.name], + "command": command, + "target": str(target.resolve()), + }) + return jobs + + +def _parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--checkpoint", required=True) + parser.add_argument( + "--expected-checkpoint-sha256", default=CHECKPOINT_SHA256, + ) + parser.add_argument("--outdir", required=True) + parser.add_argument("--gpu-indices", default="1,3") + parser.add_argument( + "--selector", + choices=("margin", "safemppi_cost", "balanced_rank"), + default="margin", + ) + parser.add_argument("--verifier-workers", type=int, default=8) + parser.add_argument("--seed", type=int, default=20260724) + parser.add_argument("--eval-ep0", type=int, default=260000) + parser.add_argument("--eval-noise-seed", type=int, default=20260723) + parser.add_argument("--idle-memory-mib", type=int, default=1024) + parser.add_argument("--idle-utilization-percent", type=int, default=5) + parser.add_argument("--dry-run", action="store_true") + return parser + + +def run(args) -> dict: + if not 1 <= int(args.verifier_workers) <= 8: + raise ValueError("--verifier-workers must be in [1,8]") + checkpoint = Path(args.checkpoint).resolve() + if not checkpoint.is_file(): + raise FileNotFoundError(checkpoint) + observed_checkpoint_sha = BASE.sha256_file(checkpoint) + if ( + args.expected_checkpoint_sha256 != CHECKPOINT_SHA256 + or observed_checkpoint_sha != CHECKPOINT_SHA256 + ): + raise RuntimeError( + f"checkpoint SHA mismatch: {observed_checkpoint_sha} != " + f"{CHECKPOINT_SHA256}" + ) + for module in (TRAINER, EVALUATOR): + if not module.is_file(): + raise FileNotFoundError(module) + outdir = _validated_output_root(args.outdir) + source = BASE.source_provenance() + arms = list(arm_grid(args.selector)) + all_gpus, processes, topology, selected = _select_exactly_two_gpus(args) + allocation = allocate_arms(arms, selected) + pools = BASE.allocate_cpu_pools(arms, int(args.verifier_workers)) + training_jobs = _phase_jobs( + args, arms, selected, allocation, pools, outdir, "training", + ) + contract = { + "version": 1, + "source": source, + "launcher_sha256": BASE.sha256_file(__file__), + "trainer_sha256": BASE.sha256_file(TRAINER), + "evaluator_sha256": BASE.sha256_file(EVALUATOR), + "checkpoint": str(checkpoint), + "checkpoint_sha256": observed_checkpoint_sha, + "scene_profile": SCENE_PROFILE, + "execution_selector": args.selector, + "rounds": ROUNDS, + "alphas": list(ALPHAS), + "exposure_epochs": list(EXPOSURE_EPOCHS), + "K": K, + "B": B, + "T": T, + "H": H, + "ell_initialization": { + "count": 50, + "multiplier": ELL_MULTIPLIER, + "representation": "stored proposal x0 at s=0.9", + }, + "cap": CAP, + "gp_lambda": GP_LAMBDA, + "batch": BATCH, + "lr": LR, + "ess_target": ESS_TARGET, + "seed": int(args.seed), + "eval_ep0": int(args.eval_ep0), + "eval_noise_seed": int(args.eval_noise_seed), + "eval_M_per_gamma": 50, + "eval_temperature": 1.0, + "verifier_workers_per_arm": int(args.verifier_workers), + "gpu_indices": [gpu.index for gpu in selected], + "gpu_uuids": [gpu.uuid for gpu in selected], + } + declaration = { + "status": "SFM_B1_OFFLINE_9ARM_DECLARED", + "created_at": _utc_now(), + "contract": contract, + "contract_sha256": _sha256_json(contract), + "all_gpus": [asdict(gpu) for gpu in all_gpus], + "compute_processes": processes, + "topology": topology, + "allocation": { + gpu.index: [arm.name for arm in allocation[gpu.uuid]] + for gpu in selected + }, + "training_jobs": [ + { + "arm": job["arm"].name, + "gpu_index": job["gpu"].index, + "gpu_uuid": job["gpu"].uuid, + "cpu_pool": job["cpu_pool"], + "command": job["command"], + "target": job["target"], + } + for job in training_jobs + ], + } + if args.dry_run: + print(json.dumps(declaration, indent=2, allow_nan=False)) + return declaration + + outdir.mkdir(parents=True) + declaration_path = outdir / "RUN_DECLARATION.json" + _write_json(declaration_path, declaration) + started = time.perf_counter() + for job in training_jobs: + job["log_path"] = str( + ( + outdir / "logs" / "training" + / f"{job['arm'].name}.log" + ).resolve() + ) + BASE._launch_pending(training_jobs, outdir / "logs" / "training") + training = { + arm.name: validate_training_arm( + outdir / "arms" / arm.name, + arm, + source_commit=source["commit"], + checkpoint_sha256=observed_checkpoint_sha, + seed=args.seed, + verifier_workers=args.verifier_workers, + ) + for arm in arms + } + training_marker = outdir / "TRAINING_COMPLETE.json" + _write_json(training_marker, { + "status": "SFM_B1_OFFLINE_9ARM_TRAINING_COMPLETE", + "finished_at": _utc_now(), + "source": source, + "declaration_sha256": BASE.sha256_file(declaration_path), + "arms": training, + }) + + # Recheck exclusivity between phases. A foreign job that appeared while + # training ran must not be silently shared with the common-bank evaluator. + _, _, _, evaluation_gpus = _select_exactly_two_gpus(args) + if [gpu.uuid for gpu in evaluation_gpus] != [ + gpu.uuid for gpu in selected + ]: + raise RuntimeError("GPU identity changed between training and evaluation") + evaluation_allocation = allocate_arms(arms, evaluation_gpus) + common_r0_dir = outdir / "evaluation" / "common_r0" + BASE._launch_pending( + [{ + "arm": PhaseName("common_r0"), + "gpu": evaluation_gpus[0], + "cpu_pool": next(iter(pools.values())), + "command": _common_r0_command(args, common_r0_dir), + "target": str(common_r0_dir.resolve()), + }], + outdir / "logs" / "evaluation_common_r0", + ) + common_r0_metrics = common_r0_dir / "raw_m50_offline_metrics.json" + if not common_r0_metrics.is_file(): + raise RuntimeError("common r0 evaluation did not produce its metrics") + with common_r0_metrics.open() as stream: + common_r0_payload = json.load(stream) + common_records = common_r0_payload.get("records", []) + if ( + common_r0_payload.get("status") != EVAL_STATUS + or len(common_records) != 1 + or int(common_records[0].get("round", -1)) != 0 + or common_records[0].get("cell", {}).get("checkpoint_sha256") + != CHECKPOINT_SHA256 + ): + raise RuntimeError("common r0 evaluation contract mismatch") + evaluation_jobs = _phase_jobs( + args, + arms, + evaluation_gpus, + evaluation_allocation, + pools, + outdir, + "evaluation", + ) + for job in evaluation_jobs: + job["log_path"] = str( + ( + outdir / "logs" / "evaluation" + / f"{job['arm'].name}.log" + ).resolve() + ) + BASE._launch_pending(evaluation_jobs, outdir / "logs" / "evaluation") + evaluations = { + arm.name: validate_evaluation( + outdir / "evaluation" / arm.name, + arm, + training[arm.name], + eval_ep0=args.eval_ep0, + eval_noise_seed=args.eval_noise_seed, + ) + for arm in arms + } + r0_cell_keys = {value["r0_cell_key"] for value in evaluations.values()} + noise_hashes = { + value["noise_bank_sha256"] for value in evaluations.values() + } + if len(r0_cell_keys) != 1 or len(noise_hashes) != 1: + raise RuntimeError( + "the nine evaluations do not share an identical r0/common bank" + ) + aggregate_result = aggregate( + evaluations, outdir / "evaluation" / "aggregate", + selector=args.selector, + ) + manifest = { + "status": "SFM_B1_OFFLINE_9ARM_DELIVERY_COMPLETE", + "finished_at": _utc_now(), + "wall_seconds": time.perf_counter() - started, + "source": source, + "contract": contract, + "declaration": str(declaration_path.resolve()), + "declaration_sha256": BASE.sha256_file(declaration_path), + "training_marker": str(training_marker.resolve()), + "training_marker_sha256": BASE.sha256_file(training_marker), + "training": training, + "evaluations": evaluations, + "common_r0_metrics": str(common_r0_metrics.resolve()), + "common_r0_metrics_sha256": BASE.sha256_file(common_r0_metrics), + "common_r0_cell_key": next(iter(r0_cell_keys)), + "common_noise_bank_sha256": next(iter(noise_hashes)), + "aggregate": aggregate_result, + } + delivery = outdir / "DELIVERY_COMPLETE.json" + _write_json(delivery, manifest) + print(json.dumps({ + "status": manifest["status"], + "wall_seconds": manifest["wall_seconds"], + "best_screening_cell": aggregate_result["best_screening_cell"], + "delivery": str(delivery.resolve()), + }, indent=2, allow_nan=False)) + return manifest + + +def main(argv=None) -> int: + run(_parser().parse_args(argv)) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/overnight_run_07_12_sfm/run_sfm_b1_offline_eval_funnel.py b/overnight_run_07_12_sfm/run_sfm_b1_offline_eval_funnel.py new file mode 100644 index 0000000..5388809 --- /dev/null +++ b/overnight_run_07_12_sfm/run_sfm_b1_offline_eval_funnel.py @@ -0,0 +1,516 @@ +#!/usr/bin/env python3 +"""Evaluate a completed selector sweep with M10 -> M50 staging. + +Training artifacts are immutable. All arm/round checkpoints first share one +raw temperature-one M10 bank. The top liveness-preserving screening cells are +then evaluated on a disjoint M50 bank. A later cross-selector job is +responsible for the final disjoint M100 confirmation. +""" +from __future__ import annotations + +import argparse +from concurrent.futures import ThreadPoolExecutor, as_completed +from datetime import datetime, timezone +import hashlib +import json +import os +from pathlib import Path +import subprocess +import sys +from typing import Any + + +HERE = Path(__file__).resolve().parent +EVALUATOR = HERE / "sfm_b1_offline_eval.py" +GAMMAS = 7 + + +def sha256_file(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as stream: + for chunk in iter(lambda: stream.read(1 << 20), b""): + digest.update(chunk) + return digest.hexdigest() + + +def sha256_json(value: Any) -> str: + encoded = json.dumps( + value, sort_keys=True, separators=(",", ":") + ).encode() + return hashlib.sha256(encoded).hexdigest() + + +def read_json(path: Path) -> dict: + if not path.is_file(): + raise FileNotFoundError(path) + with path.open() as stream: + return json.load(stream) + + +def write_json(path: Path, value: Any) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_suffix(path.suffix + ".tmp") + with temporary.open("w") as stream: + json.dump(value, stream, indent=2, allow_nan=False) + os.replace(temporary, path) + + +def validate_training(run_root: Path) -> tuple[dict, dict]: + declaration_path = run_root / "RUN_DECLARATION.json" + training_path = run_root / "TRAINING_COMPLETE.json" + declaration = read_json(declaration_path) + training = read_json(training_path) + if declaration.get("status") != "SFM_B1_OFFLINE_9ARM_DECLARED": + raise RuntimeError("invalid run declaration") + if training.get("status") != "SFM_B1_OFFLINE_9ARM_TRAINING_COMPLETE": + raise RuntimeError("training is not complete") + if training.get("declaration_sha256") != sha256_file(declaration_path): + raise RuntimeError("training marker does not authenticate declaration") + contract = declaration.get("contract", {}) + if declaration.get("contract_sha256") != sha256_json(contract): + raise RuntimeError("declaration contract digest mismatch") + checkpoint = Path(contract["checkpoint"]).resolve() + if sha256_file(checkpoint) != contract["checkpoint_sha256"]: + raise RuntimeError("pretrained checkpoint digest mismatch") + arms = training.get("arms", {}) + if len(arms) != 9: + raise RuntimeError(f"expected 9 trained arms, got {len(arms)}") + for name, arm in arms.items(): + checkpoints = arm.get("checkpoints", []) + if [row.get("round") for row in checkpoints] != list(range(11)): + raise RuntimeError(f"{name}: incomplete round checkpoints") + for row in checkpoints: + path = Path(row["path"]) + if sha256_file(path) != row["sha256"]: + raise RuntimeError(f"{name}: checkpoint digest mismatch: {path}") + return declaration, training + + +def pooled_row(arm: str, record: dict) -> dict: + pooled = record["cell"]["summary"]["pooled"] + clearance = pooled["successful_clearance"]["mean"] + time_to_goal = pooled["successful_time_to_goal"]["mean"] + return { + "arm": arm, + "round": int(record["round"]), + "checkpoint": record["cell"]["checkpoint"], + "checkpoint_sha256": record["cell"]["checkpoint_sha256"], + "SR": float(pooled["SR"]), + "CR": float(pooled["CR"]), + "timeout": float(pooled["timeout"]), + "Validity": float(pooled["Validity"]["mean"]), + "clearance": None if clearance is None else float(clearance), + "time_to_goal": ( + None if time_to_goal is None else float(time_to_goal) + ), + } + + +def safety_key(row: dict) -> tuple: + clearance = ( + -float(row["clearance"]) + if row["clearance"] is not None else float("inf") + ) + time_to_goal = ( + float(row["time_to_goal"]) + if row["time_to_goal"] is not None else float("inf") + ) + return ( + float(row["CR"]), + -float(row["Validity"]), + -float(row["SR"]), + clearance, + time_to_goal, + int(row["round"]), + str(row["arm"]), + ) + + +def fallback_key(row: dict) -> tuple: + return ( + -float(row["SR"]), + float(row["CR"]), + -float(row["Validity"]), + int(row["round"]), + str(row["arm"]), + ) + + +def choose_candidates( + rows: list[dict], r0: dict, *, top_k: int +) -> tuple[list[dict], dict]: + post = [row for row in rows if int(row["round"]) > 0] + eligible = [ + row for row in post if float(row["SR"]) >= float(r0["SR"]) + ] + ordered = sorted(eligible, key=safety_key) + fallback_used = False + if len(ordered) < top_k: + fallback_used = True + seen = { + (row["arm"], row["round"], row["checkpoint_sha256"]) + for row in ordered + } + for row in sorted(post, key=fallback_key): + key = (row["arm"], row["round"], row["checkpoint_sha256"]) + if key not in seen: + ordered.append(row) + seen.add(key) + if len(ordered) >= top_k: + break + return ordered[:top_k], { + "rule": ( + "among post-expansion cells with SR >= common-r0 SR, minimize CR, " + "then maximize window Validity and SR; if fewer than top-k pass " + "the liveness gate, supplement by highest SR" + ), + "r0_SR_gate": float(r0["SR"]), + "eligible_cells": len(eligible), + "fallback_used": fallback_used, + } + + +def evaluator_command( + checkpoints: list[str], + labels: list[str], + *, + scene_profile: str, + ep0: int, + noise_seed: int, + m_per_gamma: int, + workers: int, + cache_dir: Path, + output_dir: Path, + temperature: float = 1.0, + temperature_by_gamma: list[float] | None = None, +) -> list[str]: + command = [ + sys.executable, + str(EVALUATOR), + "--checkpoints", + *checkpoints, + "--labels", + *labels, + "--scene-profile", + scene_profile, + "--ep0", + str(ep0), + "--noise-seed", + str(noise_seed), + "--m-per-gamma", + str(m_per_gamma), + "--temperature", + str(float(temperature)), + "--device", + "cuda:0", + "--workers", + str(workers), + "--cache-dir", + str(cache_dir), + "--output-dir", + str(output_dir), + ] + if temperature_by_gamma is not None: + command.extend([ + "--temperature-by-gamma", + *map(str, temperature_by_gamma), + ]) + return command + + +def run_job( + name: str, + command: list[str], + *, + gpu_index: int, + cpu_start: int, + cpu_count: int, + log_dir: Path, +) -> dict: + log_dir.mkdir(parents=True, exist_ok=True) + log_path = log_dir / f"{name}.log" + environment = os.environ.copy() + environment["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" + environment["CUDA_VISIBLE_DEVICES"] = str(gpu_index) + environment["PYTHONPATH"] = str(HERE) + cpu_end = cpu_start + cpu_count - 1 + launched = [ + "taskset", "-c", f"{cpu_start}-{cpu_end}", *command + ] + with log_path.open("w") as stream: + completed = subprocess.run( + launched, + cwd=HERE, + env=environment, + stdout=stream, + stderr=subprocess.STDOUT, + check=False, + ) + if completed.returncode: + raise RuntimeError( + f"{name} failed with {completed.returncode}; see {log_path}" + ) + return { + "name": name, + "command": launched, + "log": str(log_path), + "log_sha256": sha256_file(log_path), + } + + +def run_parallel( + jobs: list[dict], + *, + gpu_index: int, + cpu_start: int, + cpu_count: int, + log_dir: Path, +) -> list[dict]: + results = [] + with ThreadPoolExecutor(max_workers=len(jobs)) as executor: + futures = { + executor.submit( + run_job, + job["name"], + job["command"], + gpu_index=gpu_index, + cpu_start=cpu_start + index * cpu_count, + cpu_count=cpu_count, + log_dir=log_dir, + ): job["name"] + for index, job in enumerate(jobs) + } + for future in as_completed(futures): + result = future.result() + print(f"COMPLETE {result['name']}", flush=True) + results.append(result) + return sorted(results, key=lambda row: row["name"]) + + +def run(args) -> dict: + started_at = datetime.now(timezone.utc) + run_root = Path(args.run_root).resolve() + output_root = Path(args.output_dir).resolve() + if output_root.exists(): + raise FileExistsError(output_root) + output_root.mkdir(parents=True) + declaration, training = validate_training(run_root) + contract = declaration["contract"] + selector = str(contract["execution_selector"]) + checkpoint = str(Path(contract["checkpoint"]).resolve()) + scene_profile = str(contract["scene_profile"]) + + screen_root = output_root / "screening_m10" + screen_cache = screen_root / "cache" + common_dir = screen_root / "common_r0" + common_command = evaluator_command( + [checkpoint], + ["r0"], + scene_profile=scene_profile, + ep0=args.screen_ep0, + noise_seed=args.screen_noise_seed, + m_per_gamma=args.screen_m, + workers=args.workers, + cache_dir=screen_cache, + output_dir=common_dir, + ) + common_job = run_job( + "screen_common_r0", + common_command, + gpu_index=args.gpu_index, + cpu_start=args.cpu_start, + cpu_count=args.workers, + log_dir=output_root / "logs", + ) + common_metrics = read_json( + common_dir / f"raw_m{args.screen_m}_offline_metrics.json" + ) + r0 = pooled_row("pretrained", common_metrics["records"][0]) + + screen_jobs = [] + for arm_name, arm in sorted(training["arms"].items()): + checkpoints = [checkpoint] + [ + row["path"] for row in arm["checkpoints"] if row["round"] > 0 + ] + labels = ["r0"] + [ + f"r{row['round']}" + for row in arm["checkpoints"] if row["round"] > 0 + ] + output = screen_root / arm_name + screen_jobs.append({ + "name": f"screen_{arm_name}", + "output": output, + "command": evaluator_command( + checkpoints, + labels, + scene_profile=scene_profile, + ep0=args.screen_ep0, + noise_seed=args.screen_noise_seed, + m_per_gamma=args.screen_m, + workers=args.workers, + cache_dir=screen_cache, + output_dir=output, + ), + }) + screen_logs = run_parallel( + screen_jobs, + gpu_index=args.gpu_index, + cpu_start=args.cpu_start, + cpu_count=args.workers, + log_dir=output_root / "logs", + ) + screening_rows = [] + for job in screen_jobs: + payload = read_json( + job["output"] + / f"raw_m{args.screen_m}_offline_metrics.json" + ) + screening_rows.extend( + pooled_row(job["name"].removeprefix("screen_"), record) + for record in payload["records"] + if int(record["round"]) > 0 + ) + candidates, selection = choose_candidates( + screening_rows, r0, top_k=args.top_k + ) + write_json(output_root / "SCREENING_COMPLETE.json", { + "status": "SFM_B1_OFFLINE_M10_SCREENING_COMPLETE", + "selector": selector, + "bank": { + "M_per_gamma": args.screen_m, + "ep0": args.screen_ep0, + "noise_seed": args.screen_noise_seed, + }, + "common_r0": r0, + "selection": selection, + "selected_candidates": candidates, + "rows": screening_rows, + "logs": [common_job, *screen_logs], + }) + + confirm_root = output_root / "confirmation_m50" + confirm_cache = confirm_root / "cache" + confirm_jobs = [] + for index, candidate in enumerate(candidates): + name = f"candidate_{index:02d}_{candidate['arm']}_r{candidate['round']}" + output = confirm_root / name + confirm_jobs.append({ + "name": name, + "candidate": candidate, + "output": output, + "command": evaluator_command( + [checkpoint, candidate["checkpoint"]], + ["r0", f"r{candidate['round']}"], + scene_profile=scene_profile, + ep0=args.confirm_ep0, + noise_seed=args.confirm_noise_seed, + m_per_gamma=args.confirm_m, + workers=args.workers, + cache_dir=confirm_cache, + output_dir=output, + ), + }) + confirm_logs = run_parallel( + confirm_jobs, + gpu_index=args.gpu_index, + cpu_start=args.cpu_start, + cpu_count=args.workers, + log_dir=output_root / "logs", + ) + confirmation_rows = [] + r0_confirm = None + for job in confirm_jobs: + payload = read_json( + job["output"] + / f"raw_m{args.confirm_m}_offline_metrics.json" + ) + if r0_confirm is None: + r0_confirm = pooled_row("pretrained", payload["records"][0]) + confirmation_rows.append( + pooled_row(job["candidate"]["arm"], payload["records"][1]) + ) + winner_rows, confirmation_selection = choose_candidates( + confirmation_rows, r0_confirm, top_k=1 + ) + winner = winner_rows[0] + completed_at = datetime.now(timezone.utc) + result = { + "status": "SFM_B1_OFFLINE_SELECTOR_FUNNEL_COMPLETE", + "selector": selector, + "role": ( + "M10 common-bank screening followed by disjoint M50 selector " + "confirmation; no final claim or M100 confirmation" + ), + "source_commit": subprocess.check_output( + ["git", "rev-parse", "HEAD"], cwd=HERE, text=True + ).strip(), + "training_root": str(run_root), + "training_marker_sha256": sha256_file( + run_root / "TRAINING_COMPLETE.json" + ), + "checkpoint_sha256": contract["checkpoint_sha256"], + "gpu_index": args.gpu_index, + "cpu_range": [ + args.cpu_start, + args.cpu_start + 9 * args.workers - 1, + ], + "screening": { + "bank": { + "M_per_gamma": args.screen_m, + "ep0": args.screen_ep0, + "noise_seed": args.screen_noise_seed, + }, + "common_r0": r0, + "selection": selection, + "selected_candidates": candidates, + }, + "confirmation": { + "bank": { + "M_per_gamma": args.confirm_m, + "ep0": args.confirm_ep0, + "noise_seed": args.confirm_noise_seed, + }, + "common_r0": r0_confirm, + "rows": confirmation_rows, + "selection": confirmation_selection, + "selector_winner": winner, + }, + "logs": confirm_logs, + "started_at": started_at.isoformat(), + "completed_at": completed_at.isoformat(), + "wall_seconds": (completed_at - started_at).total_seconds(), + "next_step": ( + "compare selector winners and run exactly one disjoint raw-M100 " + "confirmation on a new scenario/noise bank" + ), + } + marker = output_root / "SELECTOR_FUNNEL_COMPLETE.json" + write_json(marker, result) + print(json.dumps({ + "status": result["status"], + "selector": selector, + "winner": winner, + "marker": str(marker), + }, indent=2, allow_nan=False)) + return result + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--run-root", required=True) + parser.add_argument("--output-dir", required=True) + parser.add_argument("--gpu-index", type=int, required=True) + parser.add_argument("--cpu-start", type=int, required=True) + parser.add_argument("--workers", type=int, default=8) + parser.add_argument("--top-k", type=int, default=3) + parser.add_argument("--screen-m", type=int, default=10) + parser.add_argument("--screen-ep0", type=int, default=260_000) + parser.add_argument("--screen-noise-seed", type=int, default=2_026_072_3) + parser.add_argument("--confirm-m", type=int, default=50) + parser.add_argument("--confirm-ep0", type=int, default=270_000) + parser.add_argument("--confirm-noise-seed", type=int, default=2_026_072_4) + return parser + + +if __name__ == "__main__": + run(build_parser().parse_args()) diff --git a/overnight_run_07_12_sfm/run_sfm_b1_offline_final_confirmation.py b/overnight_run_07_12_sfm/run_sfm_b1_offline_final_confirmation.py new file mode 100644 index 0000000..ae59030 --- /dev/null +++ b/overnight_run_07_12_sfm/run_sfm_b1_offline_final_confirmation.py @@ -0,0 +1,275 @@ +#!/usr/bin/env python3 +"""Wait for three selector studies, then run one disjoint raw-M100 result.""" +from __future__ import annotations + +import argparse +import csv +from datetime import datetime, timezone +import json +from pathlib import Path +import time + +import run_sfm_b1_offline_eval_funnel as FUNNEL + + +def wait_for(path: Path, poll_seconds: int) -> dict: + while not path.is_file(): + print(f"WAITING {path}", flush=True) + time.sleep(poll_seconds) + return FUNNEL.read_json(path) + + +def margin_rows(run_root: Path) -> tuple[list[dict], dict]: + training = FUNNEL.read_json(run_root / "TRAINING_COMPLETE.json") + checkpoints = { + (name, int(row["round"])): row + for name, arm in training["arms"].items() + for row in arm["checkpoints"] + } + csv_path = ( + run_root / "evaluation" / "aggregate" + / "factorial_raw_m50_metrics.csv" + ) + rows = [] + with csv_path.open(newline="") as stream: + for source in csv.DictReader(stream): + round_index = int(source["round"]) + if round_index == 0: + continue + checkpoint = checkpoints[(source["arm"], round_index)] + rows.append({ + "arm": source["arm"], + "round": round_index, + "checkpoint": checkpoint["path"], + "checkpoint_sha256": checkpoint["sha256"], + "SR": float(source["SR"]), + "CR": float(source["CR"]), + "timeout": float(source["timeout"]), + "Validity": float(source["Validity"]), + "clearance": ( + None if not source["clearance"] + else float(source["clearance"]) + ), + "time_to_goal": ( + None if not source["time_to_goal"] + else float(source["time_to_goal"]) + ), + }) + common = FUNNEL.read_json( + run_root / "evaluation" / "common_r0" + / "raw_m50_offline_metrics.json" + ) + r0 = FUNNEL.pooled_row("pretrained", common["records"][0]) + return rows, r0 + + +def evaluate_candidates( + candidates: list[dict], + *, + checkpoint: str, + output_root: Path, + gpu_index: int, + cpu_start: int, + workers: int, + ep0: int, + noise_seed: int, +) -> tuple[list[dict], dict, list[dict]]: + cache = output_root / "cache" + jobs = [] + for index, candidate in enumerate(candidates): + name = f"margin_{index:02d}_{candidate['arm']}_r{candidate['round']}" + output = output_root / name + jobs.append({ + "name": name, + "candidate": candidate, + "output": output, + "command": FUNNEL.evaluator_command( + [checkpoint, candidate["checkpoint"]], + ["r0", f"r{candidate['round']}"], + scene_profile="double_density_velocity_ood", + ep0=ep0, + noise_seed=noise_seed, + m_per_gamma=50, + workers=workers, + cache_dir=cache, + output_dir=output, + ), + }) + logs = FUNNEL.run_parallel( + jobs, + gpu_index=gpu_index, + cpu_start=cpu_start, + cpu_count=workers, + log_dir=output_root / "logs", + ) + rows = [] + r0 = None + for job in jobs: + payload = FUNNEL.read_json( + job["output"] / "raw_m50_offline_metrics.json" + ) + if r0 is None: + r0 = FUNNEL.pooled_row("pretrained", payload["records"][0]) + rows.append( + FUNNEL.pooled_row( + job["candidate"]["arm"], payload["records"][1] + ) + ) + return rows, r0, logs + + +def run(args) -> dict: + output_root = Path(args.output_dir).resolve() + if output_root.exists(): + raise FileExistsError(output_root) + output_root.mkdir(parents=True) + margin_root = Path(args.margin_root).resolve() + cost_marker = Path(args.cost_marker).resolve() + balanced_marker = Path(args.balanced_marker).resolve() + margin_delivery = wait_for( + margin_root / "DELIVERY_COMPLETE.json", args.poll_seconds + ) + cost = wait_for(cost_marker, args.poll_seconds) + balanced = wait_for(balanced_marker, args.poll_seconds) + if margin_delivery.get("status") != ( + "SFM_B1_OFFLINE_9ARM_DELIVERY_COMPLETE" + ): + raise RuntimeError("invalid margin delivery") + for payload, selector in ( + (cost, "safemppi_cost"), + (balanced, "balanced_rank"), + ): + if ( + payload.get("status") + != "SFM_B1_OFFLINE_SELECTOR_FUNNEL_COMPLETE" + or payload.get("selector") != selector + ): + raise RuntimeError(f"invalid {selector} funnel marker") + + margin_all, margin_r0_screen = margin_rows(margin_root) + margin_candidates, margin_screen_selection = FUNNEL.choose_candidates( + margin_all, margin_r0_screen, top_k=args.top_k + ) + checkpoint = str(Path( + FUNNEL.read_json(margin_root / "RUN_DECLARATION.json") + ["contract"]["checkpoint"] + ).resolve()) + margin_confirm, shared_r0, margin_logs = evaluate_candidates( + margin_candidates, + checkpoint=checkpoint, + output_root=output_root / "margin_confirmation_m50", + gpu_index=args.gpu_index, + cpu_start=args.cpu_start, + workers=args.workers, + ep0=args.selector_ep0, + noise_seed=args.selector_noise_seed, + ) + margin_winners, margin_confirm_selection = FUNNEL.choose_candidates( + margin_confirm, shared_r0, top_k=1 + ) + selector_rows = [ + margin_winners[0], + cost["confirmation"]["selector_winner"], + balanced["confirmation"]["selector_winner"], + ] + global_winners, global_selection = FUNNEL.choose_candidates( + selector_rows, shared_r0, top_k=1 + ) + global_winner = global_winners[0] + + m100_root = output_root / "final_m100" + m100_command = FUNNEL.evaluator_command( + [checkpoint, global_winner["checkpoint"]], + ["r0", f"r{global_winner['round']}"], + scene_profile="double_density_velocity_ood", + ep0=args.final_ep0, + noise_seed=args.final_noise_seed, + m_per_gamma=100, + workers=args.workers, + cache_dir=m100_root / "cache", + output_dir=m100_root, + ) + m100_log = FUNNEL.run_job( + "final_m100", + m100_command, + gpu_index=args.gpu_index, + cpu_start=args.cpu_start, + cpu_count=args.workers, + log_dir=output_root / "logs", + ) + m100 = FUNNEL.read_json( + m100_root / "raw_m100_offline_metrics.json" + ) + r0_m100 = FUNNEL.pooled_row("pretrained", m100["records"][0]) + winner_m100 = FUNNEL.pooled_row( + global_winner["arm"], m100["records"][1] + ) + result = { + "status": "SFM_B1_OFFLINE_FINAL_M100_COMPLETE", + "role": ( + "three selectors compared on one shared disjoint M50 bank; " + "exactly one global winner confirmed on a further disjoint M100 " + "raw temperature-one bank" + ), + "source_commit": FUNNEL.subprocess.check_output( + ["git", "rev-parse", "HEAD"], cwd=FUNNEL.HERE, text=True + ).strip(), + "margin_screening": { + "bank": {"M_per_gamma": 50, "ep0": 260_000}, + "selection": margin_screen_selection, + "candidates": margin_candidates, + }, + "shared_selector_confirmation": { + "bank": { + "M_per_gamma": 50, + "ep0": args.selector_ep0, + "noise_seed": args.selector_noise_seed, + }, + "common_r0": shared_r0, + "margin_rows": margin_confirm, + "margin_selection": margin_confirm_selection, + "selector_winners": selector_rows, + "global_selection": global_selection, + "global_winner": global_winner, + }, + "final_confirmation": { + "bank": { + "M_per_gamma": 100, + "ep0": args.final_ep0, + "noise_seed": args.final_noise_seed, + }, + "pretrained_r0": r0_m100, + "winner": winner_m100, + }, + "artifacts": { + "margin_logs": margin_logs, + "m100_log": m100_log, + }, + "completed_at": datetime.now(timezone.utc).isoformat(), + } + marker = output_root / "FINAL_M100_COMPLETE.json" + FUNNEL.write_json(marker, result) + print(json.dumps(result, indent=2, allow_nan=False)) + return result + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--margin-root", required=True) + parser.add_argument("--cost-marker", required=True) + parser.add_argument("--balanced-marker", required=True) + parser.add_argument("--output-dir", required=True) + parser.add_argument("--gpu-index", type=int, default=3) + parser.add_argument("--cpu-start", type=int, default=144) + parser.add_argument("--workers", type=int, default=8) + parser.add_argument("--top-k", type=int, default=3) + parser.add_argument("--poll-seconds", type=int, default=60) + parser.add_argument("--selector-ep0", type=int, default=270_000) + parser.add_argument("--selector-noise-seed", type=int, default=2_026_072_4) + parser.add_argument("--final-ep0", type=int, default=280_000) + parser.add_argument("--final-noise-seed", type=int, default=2_026_072_5) + return parser + + +if __name__ == "__main__": + run(build_parser().parse_args()) diff --git a/overnight_run_07_12_sfm/run_sfm_b1_offline_queue.sh b/overnight_run_07_12_sfm/run_sfm_b1_offline_queue.sh new file mode 100755 index 0000000..68ada4f --- /dev/null +++ b/overnight_run_07_12_sfm/run_sfm_b1_offline_queue.sh @@ -0,0 +1,136 @@ +#!/usr/bin/env bash +set -euo pipefail + +if [[ $# -ne 3 ]]; then + echo "usage: $0 CHECKPOINT SMOKE_OUTDIR FULL_OUTDIR" >&2 + exit 2 +fi + +CHECKPOINT="$(realpath "$1")" +SMOKE_OUTDIR="$(realpath -m "$2")" +FULL_OUTDIR="$(realpath -m "$3")" +HERE="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +PYTHON="${PYTHON:-python}" +EXPECTED_SHA="1b5179c935d3eeff8824967d707d64cc9bab273949ee1f0e4f190172bab1b215" +POLL_SECONDS="${POLL_SECONDS:-20}" +IDLE_POLLS_REQUIRED="${IDLE_POLLS_REQUIRED:-3}" +GPU_INDICES="${GPU_INDICES:-1,3}" +SMOKE_GPU="${GPU_INDICES%%,*}" +EXECUTION_SELECTOR="${EXECUTION_SELECTOR:-margin}" +LOG="${QUEUE_LOG:-${FULL_OUTDIR}.queue.log}" + +mkdir -p "$(dirname "$LOG")" +exec >>"$LOG" 2>&1 + +echo "$(date -Is) QUEUE_START" +echo "source=$(git -C "$HERE/.." rev-parse HEAD)" +echo "checkpoint=$CHECKPOINT" +echo "smoke_outdir=$SMOKE_OUTDIR" +echo "full_outdir=$FULL_OUTDIR" +echo "gpu_indices=$GPU_INDICES" +echo "execution_selector=$EXECUTION_SELECTOR" + +if [[ ! -f "$CHECKPOINT" ]]; then + echo "checkpoint does not exist: $CHECKPOINT" >&2 + exit 1 +fi +if [[ "$(sha256sum "$CHECKPOINT" | awk '{print $1}')" != "$EXPECTED_SHA" ]]; then + echo "checkpoint SHA-256 mismatch" >&2 + exit 1 +fi +if [[ -e "$SMOKE_OUTDIR" || -e "$FULL_OUTDIR" ]]; then + echo "output roots must both be absent" >&2 + exit 1 +fi + +if [[ -n "${CONDA_PREFIX:-}" && -d "$CONDA_PREFIX/lib" ]]; then + export LD_LIBRARY_PATH="$CONDA_PREFIX/lib${LD_LIBRARY_PATH:+:$LD_LIBRARY_PATH}" +fi +export PYTHONPATH="$HERE${PYTHONPATH:+:$PYTHONPATH}" +export CUDA_DEVICE_ORDER=PCI_BUS_ID + +idle_polls=0 +while (( idle_polls < IDLE_POLLS_REQUIRED )); do + process_count="$( + nvidia-smi -i "$GPU_INDICES" --query-compute-apps=pid --format=csv,noheader | + sed '/^[[:space:]]*$/d' | wc -l + )" + bad_gpu_count="$( + nvidia-smi -i "$GPU_INDICES" \ + --query-gpu=memory.used,utilization.gpu \ + --format=csv,noheader,nounits | + awk -F, '{if ($1+0 > 1024 || $2+0 > 5) bad++} END {print bad+0}' + )" + gpu_count="$( + nvidia-smi -i "$GPU_INDICES" --query-gpu=index --format=csv,noheader,nounits | wc -l + )" + if [[ "$gpu_count" -eq 2 && "$process_count" -eq 0 && "$bad_gpu_count" -eq 0 ]]; then + idle_polls=$((idle_polls + 1)) + echo "$(date -Is) IDLE_CONFIRMATION ${idle_polls}/${IDLE_POLLS_REQUIRED}" + else + idle_polls=0 + echo "$(date -Is) GPU_BUSY processes=$process_count bad_gpus=$bad_gpu_count" + fi + if (( idle_polls < IDLE_POLLS_REQUIRED )); then + sleep "$POLL_SECONDS" + fi +done + +cd "$HERE" +echo "$(date -Is) SMOKE_START" +CUDA_VISIBLE_DEVICES="$SMOKE_GPU" "$PYTHON" sfm_b1_offline_exec.py \ + --checkpoint "$CHECKPOINT" \ + --outdir "$SMOKE_OUTDIR" \ + --alpha 0.01 \ + --exposure-epochs 1 \ + --selector "$EXECUTION_SELECTOR" \ + --rounds 1 \ + --verifier-workers 32 \ + --seed 20260724 \ + --device cuda:0 \ + --smoke + +"$PYTHON" - "$SMOKE_OUTDIR" <<'PY' +import json +import math +import pathlib +import sys + +root = pathlib.Path(sys.argv[1]) +with (root / "COMPLETE.json").open() as stream: + payload = json.load(stream) +assert payload["status"] == "SFM_B1_OFFLINE_EXEC_COMPLETE" +assert payload["source"]["tracked_worktree_clean"] is True +assert len(payload["history"]) == 1 +row = payload["history"][0] +summary = row["gather"]["shard"] +assert summary["D"] == summary["contexts"] +assert summary["D"] == summary["Dplus"] + summary["Dminus"] +assert summary["errors"] == summary["unresolved_contexts"] == 0 +assert row["gather"]["counts"]["B_queries"] == 4 * summary["contexts"] +assert row["replay"]["optimizer_steps"] == math.ceil(summary["D"] / 128) +assert row["replay"]["positive_total_visits"] == summary["Dplus"] +assert row["replay"]["negative_total_visits"] == summary["Dminus"] +print(json.dumps({ + "status": "SMOKE_VALIDATED", + "contexts": summary["contexts"], + "Dplus": summary["Dplus"], + "Dminus": summary["Dminus"], + "NVP": row["gather"]["counts"].get("NVP_contexts", 0), + "optimizer_steps": row["replay"]["optimizer_steps"], + "wall_seconds": row["wall_seconds"], +}, sort_keys=True)) +PY + +echo "$(date -Is) SMOKE_VALIDATED_FULL_START" +"$PYTHON" run_sfm_b1_offline_9arm.py \ + --checkpoint "$CHECKPOINT" \ + --expected-checkpoint-sha256 "$EXPECTED_SHA" \ + --outdir "$FULL_OUTDIR" \ + --gpu-indices "$GPU_INDICES" \ + --selector "$EXECUTION_SELECTOR" \ + --verifier-workers 8 \ + --seed 20260724 \ + --eval-ep0 260000 \ + --eval-noise-seed 20260723 +echo "$(date -Is) FULL_DELIVERY_COMPLETE" diff --git a/overnight_run_07_12_sfm/run_sfm_b1_r2_9arm.py b/overnight_run_07_12_sfm/run_sfm_b1_r2_9arm.py new file mode 100644 index 0000000..763b6fa --- /dev/null +++ b/overnight_run_07_12_sfm/run_sfm_b1_r2_9arm.py @@ -0,0 +1,646 @@ +#!/usr/bin/env python3 +"""Fail-closed launcher for the two-round SFM B1 alpha/replay sweep. + +The launcher deliberately knows only the narrow CLI contract of +``sfm_b1_r2_alpha_replay.py``: + + --checkpoint PATH --outdir ABSENT_DIR + --alpha FLOAT --replay-epochs INT + --verifier-workers INT --seed INT --device cuda:0 + +Each arm must atomically write ``COMPLETE.json`` with status +``R2_ALPHA_REPLAY_COMPLETE`` and authenticated ``round_00.pt`` through +``round_02.pt`` sidecars. This launcher does not import training code and +does not evaluate checkpoints. Once all arms validate, it writes a compact +``CHECKPOINT_INDEX.json`` for a separate common-bank evaluator. +""" +from __future__ import annotations + +import argparse +from dataclasses import asdict, dataclass +from datetime import datetime, timezone +import hashlib +import json +import math +import os +from pathlib import Path +import shutil +import signal +import subprocess +import sys +import time + + +HERE = Path(__file__).resolve().parent +ROOT = HERE.parent +TRAINER = HERE / "sfm_b1_r2_alpha_replay.py" +ALPHAS = (0.0, 0.01, 0.1) +REPLAY_EPOCHS = (1, 10, 100) +ROUNDS = 2 +ARM_STATUS = "R2_ALPHA_REPLAY_COMPLETE" +MAX_VERIFIER_WORKERS = 8 +ELL = 0.24210826720721101 +CAP = 256 +SCENE_PROFILE = "double_density_velocity_ood" + + +def _utc_now() -> str: + return datetime.now(timezone.utc).isoformat() + + +def sha256_file(path: str | os.PathLike[str]) -> str: + digest = hashlib.sha256() + with open(path, "rb") as stream: + for chunk in iter(lambda: stream.read(1 << 20), b""): + digest.update(chunk) + return digest.hexdigest() + + +def _sha256_json(value) -> str: + encoded = json.dumps( + value, sort_keys=True, separators=(",", ":"), allow_nan=False, + ).encode() + return hashlib.sha256(encoded).hexdigest() + + +def _write_json(path: str | os.PathLike[str], payload) -> None: + path = Path(path) + path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_name(path.name + ".tmp") + with temporary.open("w") as stream: + json.dump(payload, stream, indent=2, allow_nan=False) + os.replace(temporary, path) + + +@dataclass(frozen=True) +class Arm: + alpha: float + replay_epochs: int + + @property + def name(self) -> str: + # Match ExperimentConfig.arm_name without importing the training module. + alpha = str(float(self.alpha)).replace(".", "p") + return f"margin_alpha{alpha}_epochs{self.replay_epochs:03d}" + + +def arm_grid() -> tuple[Arm, ...]: + return tuple( + Arm(alpha, epochs) for alpha in ALPHAS for epochs in REPLAY_EPOCHS + ) + + +@dataclass(frozen=True) +class GPU: + index: str + uuid: str + name: str + memory_total_mib: int + memory_used_mib: int + utilization_percent: int + pci_bus_id: str + + +def _nvidia_lines(arguments: list[str]) -> list[str]: + try: + result = subprocess.run( + ["nvidia-smi", *arguments], check=True, text=True, + stdout=subprocess.PIPE, stderr=subprocess.PIPE, + ) + except (FileNotFoundError, subprocess.CalledProcessError) as error: + raise RuntimeError(f"nvidia-smi query failed: {error}") from error + return [line.strip() for line in result.stdout.splitlines() if line.strip()] + + +def gpu_snapshot() -> tuple[list[GPU], list[dict], str]: + rows = _nvidia_lines([ + "--query-gpu=index,uuid,name,memory.total,memory.used,utilization.gpu,pci.bus_id", + "--format=csv,noheader,nounits", + ]) + gpus = [] + for row in rows: + values = [value.strip() for value in row.split(",")] + if len(values) != 7: + raise RuntimeError(f"unexpected nvidia-smi GPU row: {row}") + gpus.append(GPU( + index=values[0], uuid=values[1], name=values[2], + memory_total_mib=int(values[3]), memory_used_mib=int(values[4]), + utilization_percent=int(values[5]), pci_bus_id=values[6], + )) + processes = [] + for row in _nvidia_lines([ + "--query-compute-apps=gpu_uuid,pid,process_name,used_memory", + "--format=csv,noheader,nounits", + ]): + values = [value.strip() for value in row.split(",")] + if len(values) == 4: + processes.append(dict( + gpu_uuid=values[0], pid=int(values[1]), process_name=values[2], + used_memory_mib=int(values[3]), + )) + try: + topology = subprocess.run( + ["nvidia-smi", "topo", "-m"], check=True, text=True, + stdout=subprocess.PIPE, stderr=subprocess.PIPE, + ).stdout + except (FileNotFoundError, subprocess.CalledProcessError): + topology = "" + return gpus, processes, topology + + +def select_idle_gpus(gpus: list[GPU], processes: list[dict], requested: str, + *, max_memory_mib: int, max_utilization: int) -> list[GPU]: + if requested == "auto": + candidates = list(gpus) + else: + indices = [value.strip() for value in requested.split(",") if value.strip()] + if len(indices) != len(set(indices)) or not indices: + raise ValueError("--gpu-indices must be 'auto' or unique comma-separated indices") + by_index = {gpu.index: gpu for gpu in gpus} + missing = [index for index in indices if index not in by_index] + if missing: + raise RuntimeError(f"requested GPU indices are unavailable: {missing}") + candidates = [by_index[index] for index in indices] + active = {row["gpu_uuid"] for row in processes} + busy = [ + gpu for gpu in candidates + if (gpu.uuid in active or gpu.memory_used_mib > int(max_memory_mib) + or gpu.utilization_percent > int(max_utilization)) + ] + if requested != "auto" and busy: + detail = [ + dict(index=gpu.index, uuid=gpu.uuid, memory_used_mib=gpu.memory_used_mib, + utilization_percent=gpu.utilization_percent, + compute_process=gpu.uuid in active) + for gpu in busy + ] + raise RuntimeError(f"explicitly requested GPUs are not idle: {detail}") + selected = [gpu for gpu in candidates if gpu not in busy] + if not selected: + raise RuntimeError("no idle GPU satisfies the launch contract") + return selected + + +def assign_arms(arms: list[Arm], gpus: list[GPU], + *, max_arms_per_gpu: int = 3) -> dict[str, list[Arm]]: + """Create a deterministic balanced allocation. + + With four GPUs and the declared grid this places all three one-step arms + together and one 10/100 pair on each remaining GPU, yielding 3/2/2/2. + For other GPU counts, a least-loaded greedy allocation is used. + """ + if not gpus: + raise ValueError("at least one GPU is required") + if int(max_arms_per_gpu) < 1: + raise ValueError("max_arms_per_gpu must be positive") + if len(arms) > len(gpus) * int(max_arms_per_gpu): + raise RuntimeError( + f"{len(arms)} arms exceed {len(gpus)} GPUs x " + f"{int(max_arms_per_gpu)} arms/GPU" + ) + ordered_gpus = sorted(gpus, key=lambda gpu: int(gpu.index)) + allocation = {gpu.uuid: [] for gpu in ordered_gpus} + if len(arms) == 9 and len(ordered_gpus) == 4 and set(arms) == set(arm_grid()): + by_steps = { + steps: sorted( + [arm for arm in arms if arm.replay_epochs == steps], + key=lambda arm: arm.alpha, + ) + for steps in REPLAY_EPOCHS + } + allocation[ordered_gpus[0].uuid].extend(by_steps[1]) + for gpu, ten, hundred in zip( + ordered_gpus[1:], by_steps[10], by_steps[100]): + allocation[gpu.uuid].extend((ten, hundred)) + return allocation + # The fallback cost is dominated by gathering, so arm count is the primary + # balance term; replay work is only a deterministic tie-break. + for arm in sorted( + arms, key=lambda value: (-value.replay_epochs, value.alpha)): + eligible = [ + gpu for gpu in ordered_gpus + if len(allocation[gpu.uuid]) < int(max_arms_per_gpu) + ] + gpu = min( + eligible, + key=lambda value: ( + len(allocation[value.uuid]), + sum(item.replay_epochs for item in allocation[value.uuid]), + int(value.index), + ), + ) + allocation[gpu.uuid].append(arm) + return allocation + + +def source_provenance() -> dict: + environment = os.environ.copy() + environment.pop("LD_LIBRARY_PATH", None) + + def git(*arguments: str) -> str: + try: + return subprocess.check_output( + ["git", *arguments], cwd=ROOT, text=True, env=environment, + ).strip() + except subprocess.CalledProcessError as error: + raise RuntimeError(f"git {' '.join(arguments)} failed") from error + + status = git("status", "--porcelain") + if status: + raise RuntimeError("source worktree must be clean before launch") + branch = git("branch", "--show-current") + if not branch: + raise RuntimeError("a named pushed branch is required") + head = git("rev-parse", "HEAD") + try: + remote = subprocess.check_output( + ["git", "ls-remote", "--heads", "origin", branch], + cwd=ROOT, text=True, env=environment, + ).split() + except subprocess.CalledProcessError as error: + raise RuntimeError("cannot authenticate origin branch") from error + if not remote or remote[0] != head: + raise RuntimeError("source HEAD is not the pushed origin branch head") + return dict(branch=branch, commit=head, remote_commit=remote[0]) + + +def _cpu_affinity() -> list[int]: + if hasattr(os, "sched_getaffinity"): + return sorted(os.sched_getaffinity(0)) + return list(range(os.cpu_count() or 1)) + + +def allocate_cpu_pools(arms: list[Arm], workers: int) -> dict[str, list[int]]: + cpus = _cpu_affinity() + needed = len(arms) * int(workers) + if len(cpus) < needed: + raise RuntimeError( + f"{len(arms)} arms x {workers} verifier workers require {needed} " + f"available CPUs, only {len(cpus)} are in the launcher affinity" + ) + return { + arm.name: cpus[index * int(workers):(index + 1) * int(workers)] + for index, arm in enumerate(arms) + } + + +def _expected_arm_contract(arm: Arm, *, checkpoint_sha256: str, scene_profile: str, + ell: float, cap: int, seed: int, + verifier_workers: int, rounds: int = ROUNDS) -> dict: + return dict( + alpha=float(arm.alpha), replay_epochs=int(arm.replay_epochs), + rounds=int(rounds), checkpoint_sha256=str(checkpoint_sha256), + scene_profile=str(scene_profile), ell=float(ell), cap=int(cap), + seed=int(seed), verifier_workers=int(verifier_workers), lr=1.0e-4, + ) + + +def validate_complete_arm(arm_dir: str | os.PathLike[str], arm: Arm, *, + checkpoint_sha256: str, scene_profile: str, + ell: float, cap: int, seed: int, + verifier_workers: int) -> dict | None: + arm_dir = Path(arm_dir) + marker = arm_dir / "COMPLETE.json" + if not marker.exists(): + if arm_dir.exists() and any(arm_dir.iterdir()): + raise RuntimeError(f"incomplete nonempty arm directory: {arm_dir}") + return None + with marker.open() as stream: + payload = json.load(stream) + if payload.get("status") != ARM_STATUS: + raise RuntimeError(f"invalid arm completion status: {marker}") + if payload.get("experiment") != arm.name: + raise RuntimeError(f"arm name mismatch in {marker}") + expected = _expected_arm_contract( + arm, checkpoint_sha256=checkpoint_sha256, + scene_profile=scene_profile, ell=ell, cap=cap, seed=seed, + verifier_workers=verifier_workers, + ) + recipe = payload.get("recipe", {}) + constants = payload.get("constants", {}) + contract = dict( + alpha=recipe.get("alpha"), replay_epochs=recipe.get("replay_epochs"), + rounds=recipe.get("rounds"), + checkpoint_sha256=payload.get("source_checkpoint_sha256"), + scene_profile=recipe.get("scene_profile"), + ell=constants.get("ell"), cap=constants.get("cap"), + seed=recipe.get("seed"), verifier_workers=recipe.get("verifier_workers"), + lr=recipe.get("lr"), + ) + if contract != expected: + raise RuntimeError( + f"arm completion contract mismatch for {arm.name}: " + f"{contract!r} != {expected!r}" + ) + history = payload.get("history") + if not isinstance(history, list) or len(history) != ROUNDS: + raise RuntimeError(f"{marker} must contain two round-history records") + history_by_round = {int(row.get("round", -1)): row for row in history} + validated = [] + for round_i in range(ROUNDS + 1): + path = (arm_dir / f"round_{round_i:02d}.pt").resolve() + expected_name = f"round_{round_i:02d}.pt" + if path.name != expected_name or not path.is_file(): + raise RuntimeError(f"missing expected checkpoint: {path}") + observed = sha256_file(path) + sidecar = Path(str(path) + ".COMPLETE.json") + if not sidecar.is_file(): + raise RuntimeError(f"missing checkpoint completion sidecar: {sidecar}") + with sidecar.open() as stream: + sidecar_payload = json.load(stream) + if (sidecar_payload.get("status") != "COMPLETE" + or sidecar_payload.get("sha256") != observed): + raise RuntimeError(f"checkpoint sidecar mismatch: {sidecar}") + if round_i > 0 and history_by_round.get(round_i, {}).get( + "checkpoint_sha256") != observed: + raise RuntimeError(f"checkpoint hash mismatch: {path}") + validated.append(dict( + round=round_i, path=str(path), sha256=observed, + complete_sidecar=str(sidecar.resolve()), + complete_sidecar_sha256=sha256_file(sidecar), + )) + return dict(marker=str(marker.resolve()), marker_sha256=sha256_file(marker), + checkpoints=validated, payload=payload) + + +def _trainer_command(args, arm: Arm, arm_dir: Path) -> list[str]: + return [ + sys.executable, str(TRAINER), + "--checkpoint", str(Path(args.checkpoint).resolve()), + "--outdir", str(arm_dir.resolve()), + "--alpha", str(arm.alpha), + "--replay-epochs", str(arm.replay_epochs), + "--verifier-workers", str(args.verifier_workers), + "--seed", str(args.seed), + "--device", "cuda:0", + ] + + +def _child_environment(gpu: GPU) -> dict[str, str]: + environment = os.environ.copy() + environment.update( + CUDA_DEVICE_ORDER="PCI_BUS_ID", + CUDA_VISIBLE_DEVICES=gpu.uuid, + OMP_NUM_THREADS="1", + MKL_NUM_THREADS="1", + OPENBLAS_NUM_THREADS="1", + NUMEXPR_NUM_THREADS="1", + TORCH_NUM_THREADS="1", + PYTHONPATH=str(HERE) + os.pathsep + environment.get("PYTHONPATH", ""), + ) + return environment + + +def _launch_pending(jobs: list[dict], log_dir: Path) -> list[str]: + taskset = shutil.which("taskset") + running = [] + log_paths = [] + log_dir.mkdir(parents=True, exist_ok=True) + try: + for job in jobs: + log_path = log_dir / f"{job['arm'].name}.log" + log_paths.append(str(log_path.resolve())) + stream = log_path.open("w") + command = list(job["command"]) + if taskset: + command = [ + taskset, "-c", ",".join(map(str, job["cpu_pool"])), *command, + ] + process = subprocess.Popen( + command, cwd=ROOT, env=_child_environment(job["gpu"]), + stdout=stream, stderr=subprocess.STDOUT, text=True, + start_new_session=True, + ) + running.append(dict( + process=process, stream=stream, log_path=str(log_path.resolve()), + arm=job["arm"], + )) + while running: + failure = None + for item in running: + code = item["process"].poll() + if code not in (None, 0): + failure = (item["arm"].name, code, item["log_path"]) + break + if failure is not None: + for item in running: + if item["process"].poll() is None: + os.killpg(item["process"].pid, signal.SIGTERM) + deadline = time.monotonic() + 10.0 + for item in running: + remaining = max(0.0, deadline - time.monotonic()) + try: + item["process"].wait(timeout=remaining) + except subprocess.TimeoutExpired: + os.killpg(item["process"].pid, signal.SIGKILL) + item["process"].wait() + raise RuntimeError( + f"arm {failure[0]} failed with code {failure[1]}; " + f"all peers were stopped; log={failure[2]}" + ) + finished = [item for item in running if item["process"].poll() == 0] + for item in finished: + item["stream"].close() + running.remove(item) + if running: + time.sleep(0.25) + except BaseException: + for item in running: + if item["process"].poll() is None: + os.killpg(item["process"].pid, signal.SIGTERM) + item["stream"].close() + raise + finally: + for item in running: + if not item["stream"].closed: + item["stream"].close() + return log_paths + + +def _parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser() + parser.add_argument("--checkpoint", required=True) + parser.add_argument("--expected-checkpoint-sha256", required=True) + parser.add_argument("--outdir", required=True) + parser.add_argument("--scene-profile", default=SCENE_PROFILE, choices=(SCENE_PROFILE,)) + parser.add_argument("--ell", type=float, default=ELL) + parser.add_argument("--cap", type=int, default=CAP) + parser.add_argument("--seed", type=int, default=20260723) + parser.add_argument("--verifier-workers", type=int, default=8) + parser.add_argument("--gpu-indices", default="auto") + parser.add_argument("--max-arms-per-gpu", type=int, default=3) + parser.add_argument("--idle-memory-mib", type=int, default=1024) + parser.add_argument("--idle-utilization-percent", type=int, default=5) + parser.add_argument("--dry-run", action="store_true") + return parser + + +def run(args) -> dict: + if not (1 <= int(args.verifier_workers) <= MAX_VERIFIER_WORKERS): + raise ValueError( + f"--verifier-workers must be in [1,{MAX_VERIFIER_WORKERS}]" + ) + if float(args.ell) != ELL or int(args.cap) != CAP: + raise ValueError(f"trainer fixes ell={ELL} and cap={CAP}") + checkpoint = Path(args.checkpoint).resolve() + if not checkpoint.is_file(): + raise FileNotFoundError(checkpoint) + observed_checkpoint_sha = sha256_file(checkpoint) + if observed_checkpoint_sha != args.expected_checkpoint_sha256: + raise RuntimeError( + f"checkpoint SHA-256 mismatch: {observed_checkpoint_sha} != " + f"{args.expected_checkpoint_sha256}" + ) + if not TRAINER.is_file(): + raise FileNotFoundError( + f"training module is not present at the frozen source path: {TRAINER}" + ) + source = source_provenance() + arms = list(arm_grid()) + outdir = Path(args.outdir).resolve() + declaration_contract = dict( + version=1, source_commit=source["commit"], + trainer_sha256=sha256_file(TRAINER), + checkpoint=str(checkpoint), checkpoint_sha256=observed_checkpoint_sha, + scene_profile=args.scene_profile, rounds=ROUNDS, + alphas=list(ALPHAS), replay_epochs=list(REPLAY_EPOCHS), + ell=float(args.ell), cap=int(args.cap), seed=int(args.seed), + verifier_workers=int(args.verifier_workers), + ) + declaration = dict( + status="SFM_B1_R2_9ARM_DECLARED", + contract=declaration_contract, + contract_sha256=_sha256_json(declaration_contract), + ) + declaration_path = outdir / "RUN_DECLARATION.json" + if declaration_path.exists(): + with declaration_path.open() as stream: + existing = json.load(stream) + if existing != declaration: + raise RuntimeError(f"existing run declaration differs: {declaration_path}") + elif outdir.exists() and any(outdir.iterdir()): + raise RuntimeError(f"nonempty output root lacks a matching declaration: {outdir}") + + completed, pending_arms = {}, [] + for arm in arms: + arm_dir = outdir / "arms" / arm.name + complete = validate_complete_arm( + arm_dir, arm, checkpoint_sha256=observed_checkpoint_sha, + scene_profile=args.scene_profile, ell=args.ell, cap=args.cap, + seed=args.seed, verifier_workers=args.verifier_workers, + ) + if complete is not None: + completed[arm.name] = complete + continue + pending_arms.append(arm) + gpus, compute_processes, topology = gpu_snapshot() + if pending_arms: + selected_gpus = select_idle_gpus( + gpus, compute_processes, args.gpu_indices, + max_memory_mib=args.idle_memory_mib, + max_utilization=args.idle_utilization_percent, + ) + allocation = assign_arms( + pending_arms, selected_gpus, max_arms_per_gpu=args.max_arms_per_gpu, + ) + by_uuid = {gpu.uuid: gpu for gpu in selected_gpus} + cpu_pools = allocate_cpu_pools(pending_arms, args.verifier_workers) + arm_gpu = { + arm: by_uuid[uuid] + for uuid, values in allocation.items() for arm in values + } + else: + selected_gpus, allocation, cpu_pools, arm_gpu = [], {}, {}, {} + pending = [ + dict( + arm=arm, gpu=arm_gpu[arm], arm_dir=outdir / "arms" / arm.name, + cpu_pool=cpu_pools[arm.name], + command=_trainer_command(args, arm, outdir / "arms" / arm.name), + ) + for arm in pending_arms + ] + plan = dict( + status="SFM_B1_R2_9ARM_DRY_RUN" if args.dry_run else "SFM_B1_R2_9ARM_PLAN", + generated_at=_utc_now(), source=source, declaration=declaration, + all_gpus=[asdict(gpu) for gpu in gpus], + selected_gpus=[asdict(gpu) for gpu in selected_gpus], + compute_processes=compute_processes, topology=topology, + allocation={ + gpu.index: [arm.name for arm in allocation[gpu.uuid]] + for gpu in selected_gpus + }, + completed_arms=sorted(completed), + pending=[dict( + arm=item["arm"].name, alpha=item["arm"].alpha, + replay_epochs=item["arm"].replay_epochs, + gpu_index=item["gpu"].index, + gpu_uuid=item["gpu"].uuid, cpu_pool=item["cpu_pool"], + command=item["command"], + ) for item in pending], + ) + if args.dry_run: + print(json.dumps(plan, indent=2, allow_nan=False)) + return plan + + outdir.mkdir(parents=True, exist_ok=True) + _write_json(declaration_path, declaration) + _write_json(outdir / "GPU_PROVENANCE.json", plan) + started = time.perf_counter() + jobs = [] + for item in pending: + item["log_path"] = str( + (outdir / "logs" / f"{item['arm'].name}.log").resolve() + ) + jobs.append(item) + logs = _launch_pending(jobs, outdir / "logs") if jobs else [] + + index_rows = [] + for arm in arms: + complete = validate_complete_arm( + outdir / "arms" / arm.name, arm, + checkpoint_sha256=observed_checkpoint_sha, + scene_profile=args.scene_profile, ell=args.ell, cap=args.cap, + seed=args.seed, verifier_workers=args.verifier_workers, + ) + if complete is None: + raise RuntimeError(f"arm returned without COMPLETE.json: {arm.name}") + index_rows.append(dict( + arm=arm.name, alpha=arm.alpha, + replay_epochs=arm.replay_epochs, + complete_marker=complete["marker"], + complete_marker_sha256=complete["marker_sha256"], + checkpoints=complete["checkpoints"], + )) + checkpoint_index = dict( + status="SFM_B1_R2_CHECKPOINT_INDEX_COMPLETE", + created_at=_utc_now(), source=source, + run_contract=declaration_contract, + common_evaluation_requirement=( + "Evaluate the unique r0 checkpoint once and every arm r1/r2 checkpoint " + "on one predeclared raw temp=1 M50/gamma common-noise bank; do not " + "force the M50 r0 estimate to equal the archival M100 statistic." + ), + arms=index_rows, + ) + index_path = outdir / "CHECKPOINT_INDEX.json" + _write_json(index_path, checkpoint_index) + complete = dict( + status="SFM_B1_R2_9ARM_TRAINING_COMPLETE", + finished_at=_utc_now(), wall_seconds=time.perf_counter() - started, + source=source, declaration_sha256=sha256_file(declaration_path), + gpu_provenance_sha256=sha256_file(outdir / "GPU_PROVENANCE.json"), + checkpoint_index=str(index_path), checkpoint_index_sha256=sha256_file(index_path), + resumed_arms=sorted(completed), launched_arms=[item["arm"].name for item in jobs], + logs=logs, + ) + _write_json(outdir / "TRAINING_COMPLETE.json", complete) + print(json.dumps(complete, indent=2, allow_nan=False)) + return complete + + +def main(argv=None) -> None: + run(_parser().parse_args(argv)) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/run_sfm_neutral_autonomous_followup.py b/overnight_run_07_12_sfm/run_sfm_neutral_autonomous_followup.py new file mode 100644 index 0000000..2136758 --- /dev/null +++ b/overnight_run_07_12_sfm/run_sfm_neutral_autonomous_followup.py @@ -0,0 +1,293 @@ +#!/usr/bin/env python3 +"""Continue the selected neutral arm to r100 only if calibrated M50 fails.""" +from __future__ import annotations + +import argparse +from datetime import datetime, timezone +import json +import os +from pathlib import Path +import subprocess +import sys +import time + +import run_sfm_neutral_gamma_temperature as GAMMA +import run_sfm_neutral_temperature_m50 as GLOBAL + + +HERE = Path(__file__).resolve().parent +STATUS = "SFM_NEUTRAL_AUTONOMOUS_FOLLOWUP_COMPLETE" +CREATIVE_TRIGGER_STATUS = "SFM_NEUTRAL_CREATIVE_SANITY_REQUIRED" + + +def _wait(path: Path, poll: int) -> dict: + while not path.is_file(): + print(f"WAITING {path}", flush=True) + time.sleep(int(poll)) + return GLOBAL._read(path) + + +def _run(command: list[str], *, log: Path, gpu: int | None = None, + cpu_range: str | None = None) -> None: + environment = os.environ.copy() + environment["PYTHONPATH"] = str(HERE) + if gpu is not None: + environment["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" + environment["CUDA_VISIBLE_DEVICES"] = str(int(gpu)) + launched = list(command) + if cpu_range is not None: + launched = ["taskset", "-c", cpu_range, *launched] + log.parent.mkdir(parents=True, exist_ok=True) + with log.open("w") as stream: + completed = subprocess.run( + launched, cwd=HERE, env=environment, + stdout=stream, stderr=subprocess.STDOUT, check=False, + ) + if completed.returncode: + raise RuntimeError(f"job failed ({completed.returncode}): {log}") + + +def run(args) -> dict: + started = datetime.now(timezone.utc) + source_commit = GLOBAL._source_gate() + gpu_provenance = GAMMA._gpu_inventory( + sorted({int(args.training_gpu), 1, 3}) + ) + gamma_delivery_path = Path(args.gamma_delivery).resolve() + gamma = _wait(gamma_delivery_path, args.poll_seconds) + gamma_delivery_sha256 = GLOBAL.FUNNEL.sha256_file(gamma_delivery_path) + output = Path(args.output_dir).resolve() + if output.exists(): + raise FileExistsError(output) + output.mkdir(parents=True) + if gamma.get("status") != GAMMA.STATUS: + raise RuntimeError("invalid calibrated M50 delivery") + if gamma.get("source_commit") != source_commit: + raise RuntimeError("calibrated M50 was produced by another source") + initial_path = Path(gamma["initial_delivery"]).resolve() + if ( + GLOBAL.FUNNEL.sha256_file(initial_path) + != gamma.get("initial_delivery_sha256") + ): + raise RuntimeError("gamma input delivery digest changed") + if gamma.get("objective_achieved") is True: + result = { + "status": STATUS, + "source_commit": source_commit, + "action": "STOP_GOAL_ACHIEVED_AT_R50_OR_EARLIER", + "gamma_delivery": str(gamma_delivery_path), + "gamma_delivery_sha256": gamma_delivery_sha256, + "gpu_provenance": gpu_provenance, + "completed_at": datetime.now(timezone.utc).isoformat(), + } + GLOBAL._source_gate(source_commit) + if GLOBAL.FUNNEL.sha256_file(gamma_delivery_path) != gamma_delivery_sha256: + raise RuntimeError("gamma delivery changed before stop decision") + GLOBAL._write(output / "DELIVERY_COMPLETE.json", result) + return result + + initial = GLOBAL._read(initial_path) + locked_best = gamma["lock"]["expanded"] + arm = str(initial["selection"]["selected_expanded"]["method"]) + training_root = Path(initial["training_root"]).resolve() + resume_root = training_root / arm + resume_delivery_path = resume_root / "DELIVERY_COMPLETE.json" + resume_delivery, resume_delivery_sha256 = GLOBAL._read_hashed( + resume_delivery_path + ) + resume_lineage = GLOBAL._authenticated_round_records(resume_delivery) + lineage_by_round = { + int(record["round"]): record for record in resume_lineage + } + cfg = resume_delivery["config"] + resume_round = int(resume_delivery["rounds"]) + if resume_round != 50: + raise RuntimeError("autonomous continuation expects an r50 source") + scenario_ids = [ + int(scenario) for record in resume_lineage + for scenario in record["scenarios"] + ] + scenario_ep0 = min(scenario_ids) + + r100_parent = output / "r100_training" + r100_root = r100_parent / arm + locked_round = int(locked_best["round"]) + if locked_round not in lineage_by_round: + raise RuntimeError("locked best is absent from the r50 lineage") + locked_lineage = lineage_by_round[locked_round] + if ( + Path(locked_best["checkpoint"]).resolve() + != Path(locked_lineage["checkpoint"]).resolve() + or locked_best["checkpoint_sha256"] + != locked_lineage["checkpoint_sha256"] + or GLOBAL.FUNNEL.sha256_file(Path(locked_best["checkpoint"])) + != locked_best["checkpoint_sha256"] + ): + raise RuntimeError("locked best is not an authenticated r50 checkpoint") + eval_rounds = sorted({0, locked_round, resume_round, *range(60, 101, 10)}) + train_command = [ + sys.executable, str(HERE / "sfm_b1_neutral_multiround.py"), + "--checkpoint", resume_delivery["checkpoint"], + "--resume-run-root", str(resume_root), + "--output-root", str(r100_root), + "--name", f"{arm}_continued_r100", + "--rounds", "100", + "--scenario-ep0", str(scenario_ep0), + "--eval-ep0", "490000", + "--eval-M", "20", + "--eval-rounds", ",".join(map(str, eval_rounds)), + "--lr", str(cfg["lr"]), + "--inner-steps", str(cfg["inner_steps"]), + "--ell", str(cfg["ell"]), + "--gp-cap", str(cfg["gp_cap"]), + "--sample-seed", str(cfg["sample_seed"]), + "--audit-seed", str(cfg["audit_seed"]), + "--train-seed", str(cfg["train_seed"]), + "--probe-seed", str(cfg["probe_seed"]), + "--noise-seed", "20260737", + "--device", "cuda:0", + "--workers", str(args.workers), + ] + if locked_round == resume_round: + resume_checkpoint = Path( + resume_root / "checkpoints" / f"round_{resume_round:02d}.pt" + ) + if ( + GLOBAL.FUNNEL.sha256_file(resume_checkpoint) + != locked_best["checkpoint_sha256"] + ): + raise RuntimeError("locked best does not match the resume anchor") + else: + train_command.extend([ + "--locked-eval-checkpoint", str(locked_best["checkpoint"]), + "--locked-eval-round", str(locked_round), + "--locked-eval-sha256", str(locked_best["checkpoint_sha256"]), + ]) + _run( + train_command, log=output / "logs" / "r50_to_r100.log", + gpu=args.training_gpu, cpu_range="16-79", + ) + r100_delivery = GLOBAL._read(r100_root / "DELIVERY_COMPLETE.json") + if int(r100_delivery.get("rounds", -1)) != 100: + raise RuntimeError("r100 continuation delivery is incomplete") + if r100_delivery.get("source", {}).get("commit") != source_commit: + raise RuntimeError("r100 continuation used another source commit") + + global_root = output / "r100_global_temperature" + global_command = [ + sys.executable, str(HERE / "run_sfm_neutral_temperature_m50.py"), "run", + "--training-root", str(r100_parent), + "--output-dir", str(global_root), + "--arm-names", arm, + "--temperatures", "0.55,0.7,0.85,1.0", + "--gpus", "1", "3", "--workers", "32", + "--screen-ep0", "490000", "--screen-M", "20", + "--validation-ep0", "500000", "--validation-M", "10", + "--validation-noise-seed", "20260738", + "--final-ep0", "510000", "--final-noise-seed", "20260739", + "--expected-final-round", "100", + "--expected-source-commit", source_commit, + ] + _run(global_command, log=output / "logs" / "r100_global_temperature.log") + + r100_gamma_root = output / "r100_gamma_temperature" + gamma_command = [ + sys.executable, str(HERE / "run_sfm_neutral_gamma_temperature.py"), + "--initial-delivery", str(global_root / "DELIVERY_COMPLETE.json"), + "--output-dir", str(r100_gamma_root), + "--temperatures", "0.55,0.7,0.85,1.0", + "--gpus", "1", "3", "--workers", "32", + "--final-ep0", "520000", "--final-noise-seed", "20260740", + "--expected-source-commit", source_commit, + ] + _run(gamma_command, log=output / "logs" / "r100_gamma_temperature.log") + r100_gamma = GLOBAL._read(r100_gamma_root / "DELIVERY_COMPLETE.json") + achieved = r100_gamma.get("objective_achieved") is True + if GLOBAL.FUNNEL.sha256_file(gamma_delivery_path) != gamma_delivery_sha256: + raise RuntimeError("r50 gamma delivery changed during continuation") + if GLOBAL.FUNNEL.sha256_file(resume_delivery_path) != resume_delivery_sha256: + raise RuntimeError("r50 resume delivery changed during continuation") + GLOBAL._source_gate(source_commit) + result = { + "status": STATUS, + "source_commit": source_commit, + "action": ( + "STOP_GOAL_ACHIEVED_AT_R100" + if achieved else "CREATIVE_SANITY_REQUIRED" + ), + "selected_arm": arm, + "locked_prior_best": locked_best, + "gpu_provenance": gpu_provenance, + "r50_training_delivery": str(resume_delivery_path), + "r50_training_delivery_sha256": resume_delivery_sha256, + "r50_gamma_delivery": str(gamma_delivery_path), + "r50_gamma_delivery_sha256": gamma_delivery_sha256, + "r100_training_delivery": str(r100_root / "DELIVERY_COMPLETE.json"), + "r100_training_delivery_sha256": GLOBAL.FUNNEL.sha256_file( + r100_root / "DELIVERY_COMPLETE.json" + ), + "r100_global_delivery": str(global_root / "DELIVERY_COMPLETE.json"), + "r100_global_delivery_sha256": GLOBAL.FUNNEL.sha256_file( + global_root / "DELIVERY_COMPLETE.json" + ), + "r100_gamma_delivery": str(r100_gamma_root / "DELIVERY_COMPLETE.json"), + "r100_gamma_delivery_sha256": GLOBAL.FUNNEL.sha256_file( + r100_gamma_root / "DELIVERY_COMPLETE.json" + ), + "ci_clean_four_metric_win": bool( + r100_gamma.get("ci_clean_four_metric_win") + ), + "objective_achieved": achieved, + "creative_sanity_priority": [ + "nontrap_progress_gated_max_margin", + "Dplus_only_if_postpositive_audit_supports_it", + "encoder_unfreeze_at_0.1x_lr", + ], + "started_at": started.isoformat(), + "completed_at": datetime.now(timezone.utc).isoformat(), + } + delivery_path = output / "DELIVERY_COMPLETE.json" + GLOBAL._write(delivery_path, result) + if not achieved: + trigger = { + "status": CREATIVE_TRIGGER_STATUS, + "action": "CREATIVE_SANITY_REQUIRED", + "source_commit": result["source_commit"], + "autonomous_delivery": str(delivery_path), + "autonomous_delivery_sha256": GLOBAL.FUNNEL.sha256_file( + delivery_path + ), + "selected_arm": arm, + "r100_training_delivery": result["r100_training_delivery"], + "r100_training_delivery_sha256": result[ + "r100_training_delivery_sha256" + ], + "r100_global_delivery": result["r100_global_delivery"], + "r100_global_delivery_sha256": result[ + "r100_global_delivery_sha256" + ], + "r100_gamma_delivery": result["r100_gamma_delivery"], + "r100_gamma_delivery_sha256": result[ + "r100_gamma_delivery_sha256" + ], + "ci_clean_four_metric_win": result["ci_clean_four_metric_win"], + "objective_achieved": result["objective_achieved"], + "created_at": datetime.now(timezone.utc).isoformat(), + } + GLOBAL._write(output / "CREATIVE_SANITY_REQUIRED.json", trigger) + print(json.dumps(result, indent=2)) + return result + + +def build_parser(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gamma-delivery", required=True) + parser.add_argument("--output-dir", required=True) + parser.add_argument("--training-gpu", type=int, default=1) + parser.add_argument("--workers", type=int, default=64) + parser.add_argument("--poll-seconds", type=int, default=60) + return parser + + +if __name__ == "__main__": + run(build_parser().parse_args()) diff --git a/overnight_run_07_12_sfm/run_sfm_neutral_creative_sanity.py b/overnight_run_07_12_sfm/run_sfm_neutral_creative_sanity.py new file mode 100644 index 0000000..7c9dbd9 --- /dev/null +++ b/overnight_run_07_12_sfm/run_sfm_neutral_creative_sanity.py @@ -0,0 +1,1673 @@ +#!/usr/bin/env python3 +"""Fail-closed causal sanity stages after the neutral r100 follow-up fails. + +The stages are intentionally sequential and never combined: + +* B: no-training, disjoint-M10 post-D+ versus post-D0 audit; +* A: five fresh rounds with only the nontrap/progress-gated margin selector; +* C: three, then at most five, fresh rounds with only E_g unfrozen at 0.1x. + +Stage B can stop the coordinator before A when it clearly identifies D0 +overwrite. Stage C is reached only after A fails its declared gate. +""" +from __future__ import annotations + +import argparse +from datetime import datetime, timezone +import hashlib +import json +import math +import os +from pathlib import Path +import subprocess +import sys +import time + +import torch + +import run_sfm_neutral_autonomous_followup as AUTO +import run_sfm_neutral_gamma_temperature as GAMMA +import run_sfm_neutral_temperature_m50 as GLOBAL +import sfm_b1_full_episode_audit as FA +import sfm_b1_neutral_multiround as TRAIN + + +HERE = Path(__file__).resolve().parent +STATUS = "SFM_NEUTRAL_CREATIVE_SANITY_COMPLETE" +STAGE_B_STATUS = "SFM_NEUTRAL_STAGE_B_POSTPHASE_AUDIT_COMPLETE" +STAGE_A_STATUS = "SFM_NEUTRAL_STAGE_A_SELECTOR_SANITY_COMPLETE" +STAGE_C_STATUS = "SFM_NEUTRAL_STAGE_C_ENCODER_SANITY_COMPLETE" +DEFAULT_AUDIT_ROUNDS = (1, 2, 5, 10, 20, 30, 40, 50, 60, 70, 80, 90, 100) + + +def _read(path: Path) -> dict: + with path.open() as stream: + return json.load(stream) + + +def _write(path: Path, payload: dict) -> None: + GLOBAL._write(path, payload) + + +def _write_final(output: Path, result: dict, source: dict) -> None: + finished = FA._source() + if ( + not finished["tracked_worktree_clean"] + or finished["commit"] != source["commit"] + ): + raise RuntimeError("creative source worktree changed during execution") + result.setdefault("source", source) + _write(output / "DELIVERY_COMPLETE.json", result) + + +def _sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as stream: + for chunk in iter(lambda: stream.read(1 << 20), b""): + digest.update(chunk) + return digest.hexdigest() + + +def _ref(path: Path) -> dict: + return {"path": str(path.resolve()), "sha256": _sha256(path.resolve())} + + +def _wait_trigger(path: Path, poll_seconds: int) -> dict: + while not path.is_file(): + print(f"WAITING {path}", flush=True) + time.sleep(int(poll_seconds)) + return _read(path) + + +def _validate_trigger( + path: Path, trigger: dict, *, expected_source: str, +) -> dict: + if ( + trigger.get("status") != AUTO.CREATIVE_TRIGGER_STATUS + or trigger.get("action") != "CREATIVE_SANITY_REQUIRED" + ): + raise RuntimeError("creative coordinator received a non-trigger marker") + if trigger.get("source_commit") != expected_source: + raise RuntimeError("creative trigger was produced by an unexpected source") + delivery_path = Path(trigger["autonomous_delivery"]).resolve() + if _sha256(delivery_path) != trigger.get("autonomous_delivery_sha256"): + raise RuntimeError("autonomous delivery digest mismatch") + delivery = _read(delivery_path) + if ( + delivery.get("status") != AUTO.STATUS + or delivery.get("action") != "CREATIVE_SANITY_REQUIRED" + or delivery.get("objective_achieved") is not False + ): + raise RuntimeError("autonomous delivery does not authorize creative sanity") + if delivery.get("source_commit") != trigger.get("source_commit"): + raise RuntimeError("trigger/autonomous source commit mismatch") + for key in ( + "selected_arm", + "r100_training_delivery", "r100_training_delivery_sha256", + "r100_global_delivery", "r100_global_delivery_sha256", + "r100_gamma_delivery", "r100_gamma_delivery_sha256", + ): + if trigger.get(key) != delivery.get(key): + raise RuntimeError(f"trigger/autonomous mismatch: {key}") + for key in ( + "r100_training_delivery", "r100_global_delivery", + "r100_gamma_delivery", + ): + referenced = Path(trigger[key]).resolve() + if _sha256(referenced) != trigger[f"{key}_sha256"]: + raise RuntimeError(f"referenced delivery digest mismatch: {key}") + gamma = _read(Path(trigger["r100_gamma_delivery"]).resolve()) + if ( + gamma.get("status") != GAMMA.STATUS + or gamma.get("objective_achieved") is not False + ): + raise RuntimeError("r100 gamma result does not require creative sanity") + return {"delivery": delivery, "gamma": gamma, "path": delivery_path} + + +def _delivery_chain(final_path: Path) -> list[tuple[Path, dict]]: + chain = [] + path = final_path.resolve() + seen = set() + while True: + if path in seen: + raise RuntimeError("resume delivery chain contains a cycle") + seen.add(path) + payload = _read(path) + if payload.get("status") != TRAIN.STATUS: + raise RuntimeError(f"invalid neutral training delivery: {path}") + chain.append((path, payload)) + resume = payload.get("resume") + if not resume: + break + path = Path(resume["delivery"]).resolve() + if _sha256(path) != resume.get("delivery_sha256"): + raise RuntimeError("resume delivery digest mismatch") + chain.reverse() + checkpoint_hashes = {payload["checkpoint_sha256"] for _, payload in chain} + if len(checkpoint_hashes) != 1: + raise RuntimeError("training chain changed pretrained checkpoint") + return chain + + +def _round_catalog(chain: list[tuple[Path, dict]]) -> dict[int, dict]: + frozen_refs = {} + for _, delivery in chain: + candidates = list(delivery.get("round_record_refs", ())) + candidates.extend(delivery.get("resume", {}).get("round_record_refs", ())) + for ref in candidates: + key = str(Path(ref["path"]).resolve()) + if key in frozen_refs and frozen_refs[key] != ref: + raise RuntimeError("conflicting frozen round reference") + frozen_refs[key] = ref + catalog = {} + for _, delivery in chain: + for marker_path in delivery.get("round_records", ()): + marker_path = Path(marker_path).resolve() + ref = frozen_refs.get(str(marker_path)) + if ref is None or _sha256(marker_path) != ref.get("sha256"): + raise RuntimeError("round marker lacks an authenticated snapshot") + marker = _read(marker_path) + round_i = int(marker.get("round", -1)) + if marker.get("status") != TRAIN.ROUND_STATUS or round_i in catalog: + raise RuntimeError("invalid or duplicate round marker") + checkpoint = Path(marker["checkpoint"]).resolve() + if ( + str(checkpoint) != str(Path(ref["post_D0"]).resolve()) + or _sha256(checkpoint) != ref.get("post_D0_sha256") + or ref.get("post_D0_sha256") != marker.get("checkpoint_sha256") + ): + raise RuntimeError("post-D0 checkpoint digest mismatch") + post_positive = ( + checkpoint.parent / f"round_{round_i:02d}_post_positive.pt" + ) + payload = torch.load( + post_positive, map_location="cpu", weights_only=False, + ) + if ( + str(post_positive) != str(Path(ref["post_Dplus"]).resolve()) + or _sha256(post_positive) != ref.get("post_Dplus_sha256") + or int(payload.get("round", -1)) != round_i + or payload.get("phase") != "post_Dplus" + ): + raise RuntimeError("post-D+ checkpoint metadata mismatch") + catalog[round_i] = { + "round": round_i, + "marker": str(marker_path), + "marker_sha256": _sha256(marker_path), + "post_positive": str(post_positive), + "post_positive_sha256": _sha256(post_positive), + "post_D0": str(checkpoint), + "post_D0_sha256": marker["checkpoint_sha256"], + "scenarios": list(map(int, marker["scenarios"])), + "record": marker, + } + rounds = sorted(catalog) + if not rounds or rounds != list(range(1, rounds[-1] + 1)): + raise RuntimeError("round lineage is not contiguous from one") + return catalog + + +def _authenticate_lock_membership(lock: dict) -> dict | None: + checkpoint = Path(lock["checkpoint"]).resolve() + if _sha256(checkpoint) != lock["checkpoint_sha256"]: + raise RuntimeError("creative lock checkpoint digest mismatch") + if lock.get("training_run") is None: + return None + run = Path(lock["training_run"]).resolve() + catalog = _round_catalog(_delivery_chain(run / "DELIVERY_COMPLETE.json")) + round_i = int(lock["round"]) + if round_i not in catalog: + raise RuntimeError("creative lock round is absent from its lineage") + frozen = catalog[round_i] + if ( + checkpoint != Path(frozen["post_D0"]).resolve() + or lock["checkpoint_sha256"] != frozen["post_D0_sha256"] + ): + raise RuntimeError("creative lock is not a frozen post-D0 checkpoint") + return frozen + + +def _bank_from(ep0: int, M: int, role: str) -> dict: + return { + "role": str(role), + "ep0": int(ep0), + "M_per_gamma": int(M), + "scenario_ids": list(range(int(ep0), int(ep0) + int(M))), + } + + +def _known_banks(chain, gamma_payload: dict) -> list[dict]: + banks = [] + for _, delivery in chain: + source = delivery.get("disjoint_raw_evaluation", {}) + if "ep0" in source and "M_per_gamma" in source: + banks.append(_bank_from( + source["ep0"], source["M_per_gamma"], "prior_training_screen", + )) + + def visit(value, path="gamma"): + if isinstance(value, dict): + if "ep0" in value and "M_per_gamma" in value: + banks.append(_bank_from( + value["ep0"], value["M_per_gamma"], path, + )) + for key, child in value.items(): + visit(child, f"{path}.{key}") + elif isinstance(value, list): + for index, child in enumerate(value): + visit(child, f"{path}[{index}]") + + visit(gamma_payload) + initial_path = gamma_payload.get("initial_delivery") + if initial_path: + visit(_read(Path(initial_path).resolve()), "gamma.initial_delivery") + unique = {} + for bank in banks: + key = (bank["ep0"], bank["M_per_gamma"]) + unique.setdefault(key, bank) + return list(unique.values()) + + +def _validate_new_banks(new_banks, known_banks, training_scenarios) -> None: + for index, bank in enumerate(new_banks): + values = set(bank["scenario_ids"]) + if values & set(training_scenarios): + raise RuntimeError(f"{bank['role']} overlaps training scenarios") + for other in [*known_banks, *new_banks[index + 1:]]: + if values & set(other["scenario_ids"]): + raise RuntimeError( + f"{bank['role']} overlaps {other['role']}" + ) + + +def _pooled(payload: dict) -> dict: + return GLOBAL._pooled(payload["records"][0]["cell"]["summary"]["pooled"]) + + +def _per_gamma(payload: dict) -> dict: + return { + gamma: GLOBAL._pooled(cell) + for gamma, cell in payload["records"][0]["cell"][ + "summary" + ]["per_gamma"].items() + } + + +def _evaluate_specs(specs, *, root, bank, noise_seed, gpus, workers): + cache = root / "cache" + jobs, lookup = [], {} + for spec in specs: + name = str(spec["name"]) + out = root / name + command = GLOBAL._raw_command( + spec["checkpoint"], round_index=int(spec["round"]), + temperature=1.0, ep0=bank["ep0"], + noise_seed=int(noise_seed), M=bank["M_per_gamma"], + workers=int(workers), output=out, cache=cache, + ) + jobs.append({"name": name, "command": command}) + lookup[name] = (spec, out) + GLOBAL._run_jobs( + jobs, gpus=list(map(int, gpus)), workers=int(workers), + log_dir=root / "logs", + ) + rows = [] + for name, (spec, out) in lookup.items(): + metrics = out / f"raw_m{bank['M_per_gamma']}_offline_metrics.json" + payload = _read(metrics) + rows.append({ + **spec, + "checkpoint_sha256": _sha256(Path(spec["checkpoint"])), + "metrics_json": str(metrics), + "metrics_sha256": _sha256(metrics), + "pooled": _pooled(payload), + "per_gamma": _per_gamma(payload), + }) + return rows + + +def _d0_gate(rows: list[dict]) -> dict: + by_round = {} + for row in rows: + by_round.setdefault(int(row["round"]), {})[row["phase"]] = row + comparisons = [] + for round_i in sorted(by_round): + phases = by_round[round_i] + if set(phases) != {"post_Dplus", "post_D0"}: + raise RuntimeError("stage-B phase pair is incomplete") + positive, neutral = phases["post_Dplus"], phases["post_D0"] + delta = { + key: neutral["pooled"][key] - positive["pooled"][key] + for key in ("SR", "CR", "timeout", "Validity") + } + liveness_harm = delta["SR"] <= -0.05 or delta["timeout"] >= 0.05 + no_material_safety_compensation = ( + delta["CR"] >= -0.03 and delta["Validity"] <= 0.03 + ) + comparisons.append({ + "round": round_i, + "post_D0_minus_post_Dplus": delta, + "clear_overwrite_cell": bool( + liveness_harm and no_material_safety_compensation + ), + }) + tail = comparisons[-min(3, len(comparisons)):] + required = max(1, math.ceil(2 * len(tail) / 3)) + clear = sum(row["clear_overwrite_cell"] for row in tail) >= required + return { + "rule": ( + "D0 clearly overwrites D+ only when >=2/3 of the final three " + "audited rounds lose >=5pp SR or gain >=5pp timeout, without " + ">3pp CR or Validity compensation" + ), + "comparisons": comparisons, + "tail_rounds": [row["round"] for row in tail], + "required_tail_cells": required, + "clear_D0_overwrite": bool(clear), + "retain_D0_unless_clear": True, + "stage_A_indicated": bool(not clear), + } + + +def _training_command( + *, checkpoint, output, name, rounds, scenario_ep0, eval_ep0, + eval_rounds, lr, inner_steps, ell, gp_cap, selector, + encoder_lr_ratio, workers, noise_seed, sample_seed, audit_seed, + train_seed, probe_seed, neutral_replay=True, resume=None, + locked_eval=None, eval_M=10, +): + command = [ + sys.executable, str(HERE / "sfm_b1_neutral_multiround.py"), + "--checkpoint", str(checkpoint), + "--output-root", str(output), + "--name", str(name), + "--rounds", str(int(rounds)), + "--scenario-ep0", str(int(scenario_ep0)), + "--eval-ep0", str(int(eval_ep0)), + "--eval-M", str(int(eval_M)), + "--eval-rounds", ",".join(map(str, eval_rounds)), + "--lr", str(float(lr)), + "--inner-steps", str(int(inner_steps)), + "--ell", str(float(ell)), + "--gp-cap", str(int(gp_cap)), + "--selector", str(selector), + "--encoder-lr-ratio", str(float(encoder_lr_ratio)), + "--noise-seed", str(int(noise_seed)), + "--sample-seed", str(int(sample_seed)), + "--audit-seed", str(int(audit_seed)), + "--train-seed", str(int(train_seed)), + "--probe-seed", str(int(probe_seed)), + "--device", "cuda:0", + "--workers", str(int(workers)), + ] + if resume is not None: + command.extend(["--resume-run-root", str(resume)]) + if locked_eval is not None: + command.extend([ + "--locked-eval-checkpoint", str(locked_eval["checkpoint"]), + "--locked-eval-round", str(int(locked_eval["round"])), + "--locked-eval-sha256", str(locked_eval["checkpoint_sha256"]), + ]) + if not neutral_replay: + command.append("--no-neutral-replay") + return command + + +def _run( + command, *, log: Path, gpu: int | None, cpu_range: str | None, +) -> None: + environment = os.environ.copy() + environment["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" + if gpu is not None: + environment["CUDA_VISIBLE_DEVICES"] = str(int(gpu)) + else: + environment.pop("CUDA_VISIBLE_DEVICES", None) + environment["PYTHONPATH"] = str(HERE) + log.parent.mkdir(parents=True, exist_ok=True) + with log.open("w") as stream: + launched = list(command) + if cpu_range is not None: + launched = ["taskset", "-c", cpu_range, *launched] + completed = subprocess.run( + launched, + cwd=HERE, env=environment, stdout=stream, + stderr=subprocess.STDOUT, check=False, + ) + if completed.returncode: + raise RuntimeError(f"job failed ({completed.returncode}): {log}") + + +def _confirm_candidate( + *, training_run: Path, arm_name: str, output: Path, ep0: int, + noise_seed: int, gpus, workers: int, selected: dict, +) -> dict: + """Calibrate temperature, then apply the existing fresh-M50 objective.""" + training_delivery = _read(training_run / "DELIVERY_COMPLETE.json") + source_commit = subprocess.check_output( + ["git", "rev-parse", "HEAD"], cwd=HERE, text=True, + ).strip() + training_screen = training_delivery["disjoint_raw_evaluation"] + input_root = output / "training_input" + prepared_arm = input_root / arm_name + prepared_arm.mkdir(parents=True) + selected_round = int(selected["round"]) + selected_sha = str(selected["checkpoint_sha256"]) + filtered_records = [ + record for record in training_screen["records"] + if int(record["round"]) in (0, selected_round) + ] + if len(filtered_records) != 2: + raise RuntimeError("confirmation input must contain r0 and one winner") + winner = next( + record for record in filtered_records + if int(record["round"]) == selected_round + ) + if winner["checkpoint_sha256"] != selected_sha: + raise RuntimeError("qualification winner checkpoint digest changed") + prepared_delivery = dict(training_delivery) + prepared_delivery["disjoint_raw_evaluation"] = { + **training_screen, + "records": filtered_records, + } + prepared_delivery["confirmation_filter"] = { + "source_delivery": str( + (training_run / "DELIVERY_COMPLETE.json").resolve() + ), + "source_delivery_sha256": _sha256( + training_run / "DELIVERY_COMPLETE.json" + ), + "selected_round": selected_round, + "selected_checkpoint": selected["checkpoint"], + "selected_checkpoint_sha256": selected_sha, + } + prepared_delivery_path = prepared_arm / "DELIVERY_COMPLETE.json" + _write(prepared_delivery_path, prepared_delivery) + global_root = output / "global_temperature" + global_command = [ + sys.executable, str(HERE / "run_sfm_neutral_temperature_m50.py"), + "run", + "--training-root", str(input_root), + "--output-dir", str(global_root), + "--arm-names", str(arm_name), + "--temperatures", "0.55,0.7,0.85,1.0", + "--gpus", *map(str, gpus), + "--workers", str(int(workers)), + "--screen-ep0", str(int(training_screen["ep0"])), + "--screen-M", str(int(training_screen["M_per_gamma"])), + "--validation-ep0", str(int(ep0)), + "--validation-M", "10", + "--validation-noise-seed", str(int(noise_seed)), + "--final-ep0", str(int(ep0) + 1_000), + "--final-noise-seed", str(int(noise_seed) + 1), + "--expected-final-round", str(int(training_delivery["rounds"])), + "--expected-source-commit", source_commit, + ] + _run( + global_command, log=output / "logs" / "global_temperature.log", + gpu=None, cpu_range=None, + ) + global_delivery = _read(global_root / "DELIVERY_COMPLETE.json") + locked = global_delivery["selection"]["selected_expanded"] + if ( + int(locked["round"]) != selected_round + or locked["checkpoint_sha256"] != selected_sha + ): + raise RuntimeError("global confirmation changed the qualified checkpoint") + gamma_root = output / "gamma_temperature" + gamma_command = [ + sys.executable, str(HERE / "run_sfm_neutral_gamma_temperature.py"), + "--initial-delivery", str(global_root / "DELIVERY_COMPLETE.json"), + "--output-dir", str(gamma_root), + "--temperatures", "0.55,0.7,0.85,1.0", + "--gpus", *map(str, gpus), + "--workers", str(int(workers)), + "--final-ep0", str(int(ep0) + 2_000), + "--final-noise-seed", str(int(noise_seed) + 2), + "--expected-source-commit", source_commit, + ] + _run( + gamma_command, log=output / "logs" / "gamma_temperature.log", + gpu=None, cpu_range=None, + ) + global_path = global_root / "DELIVERY_COMPLETE.json" + gamma_path = gamma_root / "DELIVERY_COMPLETE.json" + gamma = _read(gamma_path) + if gamma.get("status") != GAMMA.STATUS: + raise RuntimeError("creative candidate confirmation is incomplete") + return { + "prepared_training_delivery": str(prepared_delivery_path), + "prepared_training_delivery_sha256": _sha256(prepared_delivery_path), + "global_delivery": str(global_path), + "global_delivery_sha256": _sha256(global_path), + "gamma_delivery": str(gamma_path), + "gamma_delivery_sha256": _sha256(gamma_path), + "ci_clean_four_metric_win": bool( + gamma.get("ci_clean_four_metric_win") + ), + "objective_achieved": bool(gamma.get("objective_achieved")), + "final_liveness_eligible": bool( + gamma.get("final_liveness_eligible") + ), + "final_gamma_trend_eligible": bool( + gamma.get("final_gamma_trend_eligible") + ), + } + + +def _expanded_lock( + gamma_delivery: str, *, name: str, training_run: Path | None = None, +) -> dict: + path = Path(gamma_delivery).resolve() + payload = _read(path) + if payload.get("status") != GAMMA.STATUS: + raise RuntimeError("invalid gamma-temperature candidate delivery") + record = next( + row for row in payload["final_records"] + if row["method"] == "expanded" + ) + lock = { + "name": str(name), + "checkpoint": record["checkpoint"], + "checkpoint_sha256": record["checkpoint_sha256"], + "round": int(record["round"]), + "temperature_by_gamma": record["temperature_by_gamma"], + "gamma_delivery": str(path), + "gamma_delivery_sha256": _sha256(path), + "training_run": ( + None if training_run is None else str(training_run.resolve()) + ), + } + _authenticate_lock_membership(lock) + return lock + + +def _common_best_available( + locks, *, output: Path, ep0: int, noise_seed: int, gpus, workers: int, +) -> dict: + """Rank failed-objective candidates on one common M50 bank.""" + cache = output / "cache" + jobs, destinations = [], {} + for lock in locks: + _authenticate_lock_membership(lock) + name = lock["name"] + destination = output / name + jobs.append({ + "name": f"common_{name}", + "command": GLOBAL.FUNNEL.evaluator_command( + [lock["checkpoint"]], [f"r{lock['round']}"], + scene_profile="double_density_velocity_ood", + ep0=int(ep0), noise_seed=int(noise_seed), + m_per_gamma=50, workers=int(workers), + cache_dir=cache, output_dir=destination, + temperature=1.0, + temperature_by_gamma=lock["temperature_by_gamma"], + ), + }) + destinations[name] = (lock, destination) + GLOBAL._run_jobs( + jobs, gpus=[int(gpus[0])], workers=int(workers), + log_dir=output / "logs", + ) + records = [] + for name, (lock, destination) in destinations.items(): + metrics = destination / "raw_m50_offline_metrics.json" + row = GLOBAL._record_from_cell( + _read(metrics), method=name, temperature=1.0, + ) + row["temperature"] = None + row["temperature_by_gamma"] = lock["temperature_by_gamma"] + row["metrics_json"] = str(metrics) + row["metrics_sha256"] = _sha256(metrics) + if row["checkpoint_sha256"] != lock["checkpoint_sha256"]: + raise RuntimeError("common-M50 evaluator used another checkpoint") + records.append(row) + + baseline = next( + row for row in records if row["method"] == "r100_baseline" + ) + liveness = GLOBAL._liveness_contract([baseline], baseline) + for row in records: + row["liveness_eligible"] = GLOBAL._liveness_eligible(row, liveness) + row["gamma_trend_eligible"] = GLOBAL._trend_eligible(row) + + def rank(row): + value = row["pooled"] + return ( + value["CR"], -value["Validity"], -value["clearance"], + value["time_to_goal"], -value["SR"], row["method"], + ) + + eligible = [ + row for row in records + if row["liveness_eligible"] and row["gamma_trend_eligible"] + ] + selected = None if not eligible else min(eligible, key=rank) + result = { + "status": ( + "SFM_NEUTRAL_CREATIVE_COMMON_M50_COMPLETE" + if selected is not None + else "SFM_NEUTRAL_CREATIVE_COMMON_M50_NO_ELIGIBLE_POLICY" + ), + "selection_scope": ( + "best available only among policies within 5pp SR/timeout of the " + "r100 baseline and passing every gamma-trend family; then " + "safety-first lexicographic ranking, not a four-metric win claim" + ), + "single_physical_gpu_index": int(gpus[0]), + "liveness_contract": liveness, + "bank": _bank_from(ep0, 50, "creative_common_best_M50"), + "noise_seed": int(noise_seed), + "locks": locks, + "records": records, + "selected": selected, + "selected_lock": ( + None if selected is None else next( + lock for lock in locks if lock["name"] == selected["method"] + ) + ), + "completed_at": datetime.now(timezone.utc).isoformat(), + } + _write(output / "COMMON_BEST_COMPLETE.json", result) + return result + + +def _confirm_full_run( + *, training_run: Path, output: Path, ep0: int, noise_seed: int, + gpus, workers: int, +) -> dict: + delivery = _read(training_run / "DELIVERY_COMPLETE.json") + arm_name = str(delivery["config"]["name"]) + source_commit = subprocess.check_output( + ["git", "rev-parse", "HEAD"], cwd=HERE, text=True, + ).strip() + global_root = output / "global_temperature" + global_command = [ + sys.executable, str(HERE / "run_sfm_neutral_temperature_m50.py"), + "run", "--training-root", str(training_run.parent), + "--output-dir", str(global_root), "--arm-names", arm_name, + "--temperatures", "0.55,0.7,0.85,1.0", + "--gpus", *map(str, gpus), "--workers", str(int(workers)), + "--screen-ep0", str(delivery["disjoint_raw_evaluation"]["ep0"]), + "--screen-M", str(delivery["disjoint_raw_evaluation"]["M_per_gamma"]), + "--validation-ep0", str(int(ep0)), "--validation-M", "10", + "--validation-noise-seed", str(int(noise_seed)), + "--final-ep0", str(int(ep0) + 1_000), + "--final-noise-seed", str(int(noise_seed) + 1), + "--expected-final-round", str(int(delivery["rounds"])), + "--expected-source-commit", source_commit, + ] + _run( + global_command, log=output / "logs" / "global_temperature.log", + gpu=None, cpu_range=None, + ) + gamma_root = output / "gamma_temperature" + gamma_command = [ + sys.executable, str(HERE / "run_sfm_neutral_gamma_temperature.py"), + "--initial-delivery", str(global_root / "DELIVERY_COMPLETE.json"), + "--output-dir", str(gamma_root), + "--temperatures", "0.55,0.7,0.85,1.0", + "--gpus", *map(str, gpus), "--workers", str(int(workers)), + "--final-ep0", str(int(ep0) + 2_000), + "--final-noise-seed", str(int(noise_seed) + 2), + "--expected-source-commit", source_commit, + ] + _run( + gamma_command, log=output / "logs" / "gamma_temperature.log", + gpu=None, cpu_range=None, + ) + global_path = global_root / "DELIVERY_COMPLETE.json" + gamma_path = gamma_root / "DELIVERY_COMPLETE.json" + gamma = _read(gamma_path) + return { + "global_delivery": str(global_path), + "global_delivery_sha256": _sha256(global_path), + "gamma_delivery": str(gamma_path), + "gamma_delivery_sha256": _sha256(gamma_path), + "objective_achieved": bool(gamma.get("objective_achieved")), + "final_liveness_eligible": bool(gamma.get("final_liveness_eligible")), + "final_gamma_trend_eligible": bool( + gamma.get("final_gamma_trend_eligible") + ), + } + + +def _extend_common_winner_to_r100( + common: dict, *, output: Path, eval_ep0: int, confirm_ep0: int, + noise_seed: int, gpus, workers: int, +) -> dict | None: + lock = common.get("selected_lock") + if lock is None or lock["name"] == "r100_baseline": + return None + _authenticate_lock_membership(lock) + source_run = Path(lock["training_run"]).resolve() + source_delivery = _read(source_run / "DELIVERY_COMPLETE.json") + cfg = source_delivery["config"] + marker_paths = [ + ref["path"] + for ref in source_delivery.get("resume", {}).get( + "round_record_refs", () + ) + ] + marker_paths.extend(source_delivery["round_records"]) + scenario_ids = [ + int(scenario) + for marker_path in marker_paths + for scenario in _read(Path(marker_path))["scenarios"] + ] + full_parent = output / "full_r100_training" + arm_name = f"{cfg['name']}_full_r100" + full_run = full_parent / arm_name + resume_round = int(source_delivery["rounds"]) + eval_rounds = sorted({ + 0, resume_round, int(lock["round"]), *range(10, 101, 10), + }) + command = _training_command( + checkpoint=source_delivery["checkpoint"], output=full_run, + name=arm_name, rounds=100, scenario_ep0=min(scenario_ids), + eval_ep0=int(eval_ep0), eval_rounds=eval_rounds, + lr=cfg["lr"], inner_steps=cfg["inner_steps"], ell=cfg["ell"], + gp_cap=cfg["gp_cap"], selector=cfg["selector"], + encoder_lr_ratio=cfg["encoder_lr_ratio"], workers=int(workers), + noise_seed=int(noise_seed), sample_seed=cfg["sample_seed"], + audit_seed=cfg["audit_seed"], train_seed=cfg["train_seed"], + probe_seed=cfg["probe_seed"], + neutral_replay=cfg.get("neutral_replay", True), + resume=source_run, + locked_eval=(None if int(lock["round"]) == resume_round else lock), + eval_M=20, + ) + training_gpu = ( + int(gpus[-1]) if float(cfg["encoder_lr_ratio"]) > 0 else int(gpus[0]) + ) + _run( + command, log=output / "logs" / "full_r100_training.log", + gpu=training_gpu, + cpu_range="96-159" if training_gpu == int(gpus[-1]) else "16-79", + ) + full_delivery = full_run / "DELIVERY_COMPLETE.json" + confirmation = _confirm_full_run( + training_run=full_run, output=output / "full_r100_confirmation", + ep0=int(confirm_ep0), noise_seed=int(noise_seed) + 100, + gpus=gpus, workers=int(workers), + ) + return { + "source_lock": lock, + "training_gpu": training_gpu, + "training_delivery": str(full_delivery), + "training_delivery_sha256": _sha256(full_delivery), + "confirmation": confirmation, + "objective_achieved": confirmation["objective_achieved"], + } + + +def _delivery_eval_rows(delivery: dict) -> list[dict]: + source = _read(Path(delivery["disjoint_raw_evaluation"]["file"])) + rows = [] + for row in source["records"]: + rows.append({ + "name": row["name"], + "round": int(row["round"]), + "phase": row["phase"], + "checkpoint": row["checkpoint"], + "checkpoint_sha256": row["checkpoint_sha256"], + "pooled": GLOBAL._pooled(row["cell"]["summary"]["pooled"]), + "per_gamma": { + gamma: GLOBAL._pooled(cell) + for gamma, cell in row["cell"]["summary"][ + "per_gamma" + ].items() + }, + }) + return rows + + +def _trend_ok(row: dict) -> bool: + return GLOBAL._trend_eligible(row) + + +def _selector_gate( + candidate_rows, controls, r100_reference=None, *, diagnostic_rounds=None, +) -> dict: + control_by_round = {int(row["round"]): row for row in controls} + candidate_by_round = { + int(row["round"]): row for row in candidate_rows + } + decisions = [] + for row in candidate_rows: + if int(row["round"]) == 0: + continue + control = control_by_round[int(row["round"])] + value, base = row["pooled"], control["pooled"] + safe = ( + value["CR"] <= base["CR"] + .03 + and value["Validity"] >= base["Validity"] - .03 + ) + live = ( + value["SR"] >= base["SR"] + .05 + or value["timeout"] <= base["timeout"] - .05 + ) + reference = None if r100_reference is None else r100_reference["pooled"] + noninferior_to_r100 = None if reference is None else ( + value["SR"] >= reference["SR"] - .03 + and value["timeout"] <= reference["timeout"] + .03 + and value["CR"] <= reference["CR"] + .03 + and value["Validity"] >= reference["Validity"] - .03 + ) + diagnostic_eligible = ( + diagnostic_rounds is None + or int(row["round"]) in set(map(int, diagnostic_rounds)) + ) + decisions.append({ + "round": int(row["round"]), + "checkpoint": row["checkpoint"], + "checkpoint_sha256": row["checkpoint_sha256"], + "pooled": value, + "safety_noninferior_to_matched_margin": bool(safe), + "liveness_improved_over_matched_margin": bool(live), + "noninferior_to_r100_reference": bool(noninferior_to_r100), + "diagnostic_round_eligible": bool(diagnostic_eligible), + "gamma_trend_pass": _trend_ok(row), + "eligible": bool( + safe and live and diagnostic_eligible and _trend_ok(row) + ), + }) + eligible = [row for row in decisions if row["eligible"]] + selected = None if not eligible else min( + eligible, + key=lambda item: ( + -candidate_by_round[item["round"]]["pooled"]["SR"], + candidate_by_round[item["round"]]["pooled"]["CR"], + item["round"], + ), + ) + return { + "rule": ( + "promote only if SR improves >=5pp or timeout falls >=5pp vs " + "the matched equal-dose margin round, while CR/Validity are within " + "3pp and every gamma-trend family is >=0.75; r100 is diagnostic " + "only until the mechanism reaches an equal dose" + ), + "r100_reference": r100_reference, + "decisions": decisions, + "passed": selected is not None, + "selected": selected, + } + + +def _control_gp_fields(marker: dict) -> dict: + gather = marker["gather"] + if "acquisition" in gather: + return { + "uplift": float(gather["acquisition"]["uplift"]), + "effective_rank": ( + None if "gp_diagnostics" not in gather else float( + gather["gp_diagnostics"]["kernel_effective_rank"] + ) + ), + "source": "round_marker", + } + trace_path = Path(gather["trace_path"]).resolve() + if _sha256(trace_path) != gather.get("trace_sha256"): + raise RuntimeError("legacy control gather trace digest mismatch") + trace = torch.load(trace_path, map_location="cpu", weights_only=False) + if trace.get("status") != TRAIN.RA.STATUS: + raise RuntimeError("invalid legacy control gather trace") + acquisition = trace.get("protocol", {}).get("acquisition") + if acquisition is None or "uplift" not in acquisition: + raise RuntimeError("legacy control trace lacks acquisition diagnostics") + return { + "uplift": float(acquisition["uplift"]), + "effective_rank": None, + "source": "authenticated_legacy_trace", + } + + +def _encoder_stage_gate( + delivery, candidate_rows, control_rows, control_markers, +) -> dict: + controls = {int(row["round"]): row for row in control_rows} + markers = [_read(Path(path)) for path in delivery["round_records"]] + diagnostics = [] + for marker in markers: + round_i = int(marker["round"]) + candidate = next( + row for row in candidate_rows if int(row["round"]) == round_i + ) + control = controls[round_i] + encoder = marker["encoder_diagnostics"] + cumulative = encoder["cumulative_from_reference"] + gather = marker["gather"] + control_marker = control_markers[round_i] + control_gp = _control_gp_fields(control_marker) + candidate_rank = float( + gather["gp_diagnostics"]["kernel_effective_rank"] + ) + control_rank = control_gp["effective_rank"] + unsafe = ( + cumulative["token_cosine"] < .98 + or marker["paired_trigger_probe"]["Dplus_increment"][ + "Dplus_regressed" + ] > 0 + or candidate["pooled"]["CR"] - control["pooled"]["CR"] > .03 + ) + signal = ( + candidate["pooled"]["Validity"] + - control["pooled"]["Validity"] >= .01 + or gather["acquisition"]["uplift"] - control_gp["uplift"] >= .002 + or ( + control_rank is not None + and candidate_rank - control_rank >= 1.0 + ) + ) + diagnostics.append({ + "round": round_i, + "token_cosine": encoder["token_cosine"], + "encoder_relative_drift": encoder["relative_parameter_drift"], + "cumulative_token_cosine": cumulative["token_cosine"], + "cumulative_token_rms_change": cumulative["token_rms_change"], + "cumulative_encoder_relative_drift": cumulative[ + "relative_parameter_drift" + ], + "encoder_gradient_norm_Dplus": marker["updates"]["Dplus"][ + "encoder_gradient_norms" + ], + "encoder_gradient_norm_D0": marker["updates"]["D0"][ + "encoder_gradient_norms" + ], + "Dplus_regressed": marker["paired_trigger_probe"][ + "Dplus_increment" + ]["Dplus_regressed"], + "gp_effective_rank": gather["gp_diagnostics"][ + "kernel_effective_rank" + ], + "uncertainty_uplift": gather["acquisition"]["uplift"], + "frozen_gp_effective_rank": control_rank, + "frozen_uncertainty_uplift": control_gp["uplift"], + "frozen_gp_diagnostic_source": control_gp["source"], + "delta_gp_effective_rank_vs_frozen": ( + None if control_rank is None else candidate_rank - control_rank + ), + "delta_uncertainty_uplift_vs_frozen": ( + gather["acquisition"]["uplift"] + - control_gp["uplift"] + ), + "delta_CR_vs_frozen": ( + candidate["pooled"]["CR"] - control["pooled"]["CR"] + ), + "delta_Validity_vs_frozen": ( + candidate["pooled"]["Validity"] + - control["pooled"]["Validity"] + ), + "unsafe": bool(unsafe), + "signal": bool(signal), + "eligible": bool(not unsafe and signal), + }) + eligible_rounds = [row["round"] for row in diagnostics if row["eligible"]] + return { + "rule": ( + "stop on cumulative E_g token cosine <0.98, any D+ regression, or CR >3pp " + "above matched frozen control; continue only with >=1pp Validity " + "or matched GP uplift/rank improvement" + ), + "diagnostics": diagnostics, + "unsafe": bool(not eligible_rounds), + "signal": bool(any(row["signal"] for row in diagnostics)), + "eligible_rounds": eligible_rounds, + "continue_to_round5": bool( + diagnostics and diagnostics[-1]["eligible"] + ), + } + + +def _base_specs(catalog, rounds, prefix): + return [{ + "name": f"{prefix}_r{round_i}", + "phase": "post_D0", + "round": int(round_i), + "checkpoint": catalog[int(round_i)]["post_D0"], + } for round_i in rounds] + + +def run(args) -> dict: + started = datetime.now(timezone.utc) + trigger_path = Path(args.trigger).resolve() + trigger = _wait_trigger(trigger_path, args.poll_seconds) + authenticated = _validate_trigger( + trigger_path, trigger, + expected_source=args.expected_trigger_source, + ) + source = FA._source() + if not source["tracked_worktree_clean"]: + raise RuntimeError("creative sanity requires a clean frozen worktree") + if source["commit"] != args.expected_trigger_source: + raise RuntimeError("creative sanity source does not match its trigger") + output = Path(args.output_dir).resolve() + if output.exists(): + raise FileExistsError(output) + + chain = _delivery_chain(Path(trigger["r100_training_delivery"])) + catalog = _round_catalog(chain) + selected_arm = str(trigger["selected_arm"]) + final_delivery = chain[-1][1] + if chain[0][1]["config"]["name"] != selected_arm: + raise RuntimeError("selected arm does not match r100 lineage") + audit_rounds = tuple(sorted({ + int(value) for value in str(args.audit_rounds).split(",") if value + })) + if not audit_rounds or audit_rounds[-1] != max(catalog): + raise ValueError("audit rounds must include the final lineage round") + if any(value not in catalog for value in audit_rounds): + raise ValueError("audit round is absent from the lineage") + + training_scenarios = { + scenario + for record in catalog.values() for scenario in record["scenarios"] + } + banks = [ + _bank_from(args.stage_b_ep0, 10, "stage_B_disjoint_M10"), + _bank_from(args.stage_a_ep0, 10, "stage_A_disjoint_M10"), + _bank_from(args.stage_c3_ep0, 10, "stage_C3_disjoint_M10"), + _bank_from(args.stage_c5_ep0, 10, "stage_C5_disjoint_M10"), + _bank_from( + args.stage_b_causal_ep0, 10, + "stage_B_Dplus_only_disjoint_M10", + ), + ] + confirmation_banks = [] + for label, ep0 in ( + ("B_Dplus_only", args.stage_b_confirm_ep0), + ("A_selector", args.stage_a_confirm_ep0), + ("C_encoder", args.stage_c_confirm_ep0), + ): + confirmation_banks.extend([ + _bank_from(ep0, 10, f"{label}_temperature_validation_M10"), + _bank_from(ep0 + 1_000, 50, f"{label}_global_M50"), + _bank_from(ep0 + 2_000, 50, f"{label}_gamma_fresh_M50"), + ]) + confirmation_banks.append(_bank_from( + args.common_best_ep0, 50, "creative_common_best_M50", + )) + confirmation_banks.extend([ + _bank_from(args.full_eval_ep0, 20, "creative_full_r100_screen_M20"), + _bank_from(args.full_confirm_ep0, 10, "creative_full_temperature_M10"), + _bank_from( + args.full_confirm_ep0 + 1_000, 50, + "creative_full_global_M50", + ), + _bank_from( + args.full_confirm_ep0 + 2_000, 50, + "creative_full_gamma_fresh_M50", + ), + ]) + known_banks = _known_banks(chain, authenticated["gamma"]) + _validate_new_banks( + [*banks, *confirmation_banks], known_banks, training_scenarios, + ) + output.mkdir(parents=True) + baseline_lock = _expanded_lock( + trigger["r100_gamma_delivery"], name="r100_baseline", + training_run=Path(trigger["r100_training_delivery"]).resolve().parent, + ) + creative_locks = [] + provenance = { + "status": "SFM_NEUTRAL_CREATIVE_SANITY_PREREGISTERED", + "source": source, + "trigger": str(trigger_path), + "trigger_sha256": _sha256(trigger_path), + "autonomous_delivery": str(authenticated["path"]), + "autonomous_delivery_sha256": _sha256(authenticated["path"]), + "selected_arm": selected_arm, + "lineage_catalog": { + str(round_i): { + key: value for key, value in item.items() if key != "record" + } + for round_i, item in catalog.items() + }, + "pretrained_checkpoint": final_delivery["checkpoint"], + "pretrained_checkpoint_sha256": final_delivery["checkpoint_sha256"], + "r100_training_delivery": trigger["r100_training_delivery"], + "r100_training_delivery_sha256": _sha256( + Path(trigger["r100_training_delivery"]) + ), + "audit_rounds": list(audit_rounds), + "known_banks": known_banks, + "new_banks": [*banks, *confirmation_banks], + "no_combined_arm": True, + "stage_order": [ + "B_postphase_audit", + "B_Dplus_only_if_indicated", + "A_selector", + "C_encoder", + ], + "created_at": datetime.now(timezone.utc).isoformat(), + } + _write(output / "PREREGISTRATION.json", provenance) + + stage_b_root = output / "stage_B_postphase_audit" + stage_b_specs = [] + for round_i in audit_rounds: + item = catalog[round_i] + stage_b_specs.extend([ + { + "name": f"r{round_i}_post_Dplus", "round": round_i, + "phase": "post_Dplus", "checkpoint": item["post_positive"], + }, + { + "name": f"r{round_i}_post_D0", "round": round_i, + "phase": "post_D0", "checkpoint": item["post_D0"], + }, + ]) + stage_b_rows = _evaluate_specs( + stage_b_specs, root=stage_b_root, bank=banks[0], + noise_seed=args.stage_b_noise_seed, gpus=[args.gpus[0]], + workers=args.workers, + ) + stage_b_gate = _d0_gate(stage_b_rows) + stage_b = { + "status": STAGE_B_STATUS, + "zero_new_training": True, + "bank": banks[0], + "noise_seed": int(args.stage_b_noise_seed), + "rows": stage_b_rows, + "gate": stage_b_gate, + } + _write(stage_b_root / "STAGE_COMPLETE.json", stage_b) + cfg = final_delivery["config"] + scenario_ep0 = min(training_scenarios) + stage_b_causal = None + if stage_b_gate["clear_D0_overwrite"]: + dplus_root = output / "stage_B_Dplus_only" + command = _training_command( + checkpoint=final_delivery["checkpoint"], output=dplus_root, + name=f"{selected_arm}_Dplus_only", rounds=5, + scenario_ep0=scenario_ep0, eval_ep0=banks[4]["ep0"], + eval_rounds=(0, 1, 2, 3, 4, 5), lr=cfg["lr"], + inner_steps=cfg["inner_steps"], ell=cfg["ell"], + gp_cap=cfg["gp_cap"], selector="margin", + encoder_lr_ratio=0.0, neutral_replay=False, + workers=args.workers, noise_seed=args.stage_b_causal_noise_seed, + sample_seed=cfg["sample_seed"], audit_seed=cfg["audit_seed"], + train_seed=cfg["train_seed"], probe_seed=cfg["probe_seed"], + ) + _run( + command, log=output / "logs" / "stage_B_Dplus_only.log", + gpu=args.gpus[0], cpu_range="16-79", + ) + dplus_delivery = _read(dplus_root / "DELIVERY_COMPLETE.json") + if dplus_delivery["config"].get("neutral_replay") is not False: + raise RuntimeError("stage B causal arm replayed D0") + dplus_rows = _delivery_eval_rows(dplus_delivery) + dplus_controls = _evaluate_specs( + _base_specs(catalog, range(1, 6), "margin_control"), + root=dplus_root / "matched_margin_control", bank=banks[4], + noise_seed=args.stage_b_causal_noise_seed, + gpus=[args.gpus[0]], workers=args.workers, + ) + dplus_r100 = _evaluate_specs([{ + "name": "r100_reference", "round": 100, + "phase": "post_D0", "checkpoint": catalog[100]["post_D0"], + }], root=dplus_root / "r100_reference", bank=banks[4], + noise_seed=args.stage_b_causal_noise_seed, + gpus=[args.gpus[0]], workers=args.workers)[0] + dplus_gate = _selector_gate( + dplus_rows, dplus_controls, dplus_r100, + ) + stage_b_causal = { + "status": "SFM_NEUTRAL_STAGE_B_DPLUS_ONLY_COMPLETE", + "single_change": "collect/audit D0 but omit D0 replay", + "training_delivery": _ref( + dplus_root / "DELIVERY_COMPLETE.json" + ), + "bank": banks[4], + "rows": dplus_rows, + "matched_margin_controls": dplus_controls, + "qualification_gate": dplus_gate, + "confirmation": None, + } + if dplus_gate["passed"]: + stage_b_causal["confirmation"] = _confirm_candidate( + training_run=dplus_root, + arm_name=dplus_delivery["config"]["name"], + output=dplus_root / "confirmation", + ep0=args.stage_b_confirm_ep0, + noise_seed=args.stage_b_confirm_noise_seed, + gpus=args.gpus, workers=args.workers, + selected=dplus_gate["selected"], + ) + creative_locks.append(_expanded_lock( + stage_b_causal["confirmation"]["gamma_delivery"], + name="Dplus_only", + training_run=dplus_root, + )) + _write(dplus_root / "STAGE_COMPLETE.json", stage_b_causal) + if ( + stage_b_causal["confirmation"] is not None + and stage_b_causal["confirmation"]["objective_achieved"] + ): + result = { + "status": STATUS, + "action": "STOP_GOAL_ACHIEVED_BY_DPLUS_ONLY", + "objective_achieved": True, + "stages_completed": ["B", "B_Dplus_only"], + "stage_B": _ref(stage_b_root / "STAGE_COMPLETE.json"), + "stage_B_Dplus_only": _ref( + dplus_root / "STAGE_COMPLETE.json" + ), + "completed_at": datetime.now(timezone.utc).isoformat(), + } + _write_final(output, result, source) + return result + + stage_a_root = output / "stage_A_progress_selector" + command = _training_command( + checkpoint=final_delivery["checkpoint"], output=stage_a_root, + name=f"{selected_arm}_progress_gated", rounds=5, + scenario_ep0=scenario_ep0, eval_ep0=banks[1]["ep0"], + eval_rounds=(0, 1, 2, 3, 4, 5), lr=cfg["lr"], + inner_steps=cfg["inner_steps"], ell=cfg["ell"], + gp_cap=cfg["gp_cap"], selector="progress_gated_margin", + encoder_lr_ratio=0.0, workers=args.workers, + noise_seed=args.stage_a_noise_seed, + sample_seed=cfg["sample_seed"], audit_seed=cfg["audit_seed"], + train_seed=cfg["train_seed"], probe_seed=cfg["probe_seed"], + ) + _run( + command, log=output / "logs" / "stage_A.log", + gpu=args.gpus[0], cpu_range="16-79", + ) + stage_a_delivery = _read(stage_a_root / "DELIVERY_COMPLETE.json") + if ( + stage_a_delivery["config"]["selector"] != "progress_gated_margin" + or stage_a_delivery["config"]["encoder_lr_ratio"] != 0.0 + ): + raise RuntimeError("stage A combined more than the selector change") + stage_a_rows = _delivery_eval_rows(stage_a_delivery) + controls_a = _evaluate_specs( + _base_specs(catalog, range(1, 6), "margin_control"), + root=stage_a_root / "matched_margin_control", + bank=banks[1], noise_seed=args.stage_a_noise_seed, + gpus=[args.gpus[0]], workers=args.workers, + ) + r100_a = _evaluate_specs([ + { + "name": "r100_reference", "round": 100, + "phase": "post_D0", "checkpoint": catalog[100]["post_D0"], + } + ], root=stage_a_root / "r100_reference", bank=banks[1], + noise_seed=args.stage_a_noise_seed, gpus=[args.gpus[0]], + workers=args.workers)[0] + stage_a_gate = _selector_gate(stage_a_rows, controls_a, r100_a) + stage_a = { + "status": STAGE_A_STATUS, + "single_change": "progress_gated_margin selector", + "training_delivery": _ref(stage_a_root / "DELIVERY_COMPLETE.json"), + "bank": banks[1], + "rows": stage_a_rows, + "matched_margin_controls": controls_a, + "qualification_gate": stage_a_gate, + "confirmation": None, + } + if stage_a_gate["passed"]: + stage_a["confirmation"] = _confirm_candidate( + training_run=stage_a_root, + arm_name=stage_a_delivery["config"]["name"], + output=stage_a_root / "confirmation", + ep0=args.stage_a_confirm_ep0, + noise_seed=args.stage_a_confirm_noise_seed, + gpus=args.gpus, workers=args.workers, + selected=stage_a_gate["selected"], + ) + creative_locks.append(_expanded_lock( + stage_a["confirmation"]["gamma_delivery"], + name="progress_gated_margin", + training_run=stage_a_root, + )) + _write(stage_a_root / "STAGE_COMPLETE.json", stage_a) + if ( + stage_a["confirmation"] is not None + and stage_a["confirmation"]["objective_achieved"] + ): + result = { + "status": STATUS, + "action": "STOP_GOAL_ACHIEVED_BY_STAGE_A", + "objective_achieved": True, + "stages_completed": ["B", "A"], + "stage_B": _ref(stage_b_root / "STAGE_COMPLETE.json"), + "stage_B_Dplus_only": ( + None if stage_b_causal is None else _ref( + output / "stage_B_Dplus_only" / "STAGE_COMPLETE.json" + ) + ), + "stage_A": _ref(stage_a_root / "STAGE_COMPLETE.json"), + "completed_at": datetime.now(timezone.utc).isoformat(), + } + _write_final(output, result, source) + return result + + stage_c3_root = output / "stage_C_encoder_unfreeze_r3" + command = _training_command( + checkpoint=final_delivery["checkpoint"], output=stage_c3_root, + name=f"{selected_arm}_encoder01x", rounds=3, + scenario_ep0=scenario_ep0, eval_ep0=banks[2]["ep0"], + eval_rounds=(0, 1, 2, 3), lr=cfg["lr"], + inner_steps=cfg["inner_steps"], ell=cfg["ell"], + gp_cap=cfg["gp_cap"], selector="margin", + encoder_lr_ratio=0.1, workers=args.workers, + noise_seed=args.stage_c3_noise_seed, + sample_seed=cfg["sample_seed"], audit_seed=cfg["audit_seed"], + train_seed=cfg["train_seed"], probe_seed=cfg["probe_seed"], + ) + _run( + command, log=output / "logs" / "stage_C3.log", + gpu=args.gpus[-1], cpu_range="96-159", + ) + c3_delivery = _read(stage_c3_root / "DELIVERY_COMPLETE.json") + if ( + c3_delivery["config"]["selector"] != "margin" + or c3_delivery["config"]["encoder_lr_ratio"] != 0.1 + ): + raise RuntimeError("stage C combined more than encoder unfreezing") + c3_rows = _delivery_eval_rows(c3_delivery) + c3_controls = _evaluate_specs( + _base_specs(catalog, range(1, 4), "frozen_control"), + root=stage_c3_root / "matched_frozen_control", bank=banks[2], + noise_seed=args.stage_c3_noise_seed, gpus=[args.gpus[-1]], + workers=args.workers, + ) + c3_r100 = _evaluate_specs([{ + "name": "r100_reference", "round": 100, + "phase": "post_D0", "checkpoint": catalog[100]["post_D0"], + }], root=stage_c3_root / "r100_reference", bank=banks[2], + noise_seed=args.stage_c3_noise_seed, gpus=[args.gpus[-1]], + workers=args.workers)[0] + c3_gate = _encoder_stage_gate( + c3_delivery, c3_rows, c3_controls, + {round_i: catalog[round_i]["record"] for round_i in range(1, 4)}, + ) + selector_c3 = _selector_gate( + c3_rows, c3_controls, c3_r100, + diagnostic_rounds=c3_gate["eligible_rounds"], + ) + c3 = { + "status": STAGE_C_STATUS, + "phase": "rounds_1_to_3", + "single_change": "enc_grid lr = 0.1 * trunk lr", + "training_delivery": _ref(stage_c3_root / "DELIVERY_COMPLETE.json"), + "bank": banks[2], + "rows": c3_rows, + "matched_frozen_controls": c3_controls, + "gate": c3_gate, + "promotion_gate": selector_c3, + "confirmation": None, + } + if selector_c3["passed"]: + c3["confirmation"] = _confirm_candidate( + training_run=stage_c3_root, + arm_name=c3_delivery["config"]["name"], + output=stage_c3_root / "confirmation", + ep0=args.stage_c_confirm_ep0, + noise_seed=args.stage_c_confirm_noise_seed, + gpus=args.gpus, workers=args.workers, + selected=selector_c3["selected"], + ) + creative_locks.append(_expanded_lock( + c3["confirmation"]["gamma_delivery"], + name="encoder_unfreeze_01x_r3", + training_run=stage_c3_root, + )) + _write(stage_c3_root / "STAGE_COMPLETE.json", c3) + if ( + c3["confirmation"] is not None + and c3["confirmation"]["objective_achieved"] + ): + result = { + "status": STATUS, + "action": "STOP_GOAL_ACHIEVED_BY_STAGE_C3", + "objective_achieved": True, + "stages_completed": ["B", "A", "C3"], + "stage_B": _ref(stage_b_root / "STAGE_COMPLETE.json"), + "stage_B_Dplus_only": ( + None if stage_b_causal is None else _ref( + output / "stage_B_Dplus_only" / "STAGE_COMPLETE.json" + ) + ), + "stage_A": _ref(stage_a_root / "STAGE_COMPLETE.json"), + "stage_C3": _ref(stage_c3_root / "STAGE_COMPLETE.json"), + "completed_at": datetime.now(timezone.utc).isoformat(), + } + _write_final(output, result, source) + return result + if not c3_gate["continue_to_round5"]: + common = _common_best_available( + [baseline_lock, *creative_locks], + output=output / "common_best_M50", + ep0=args.common_best_ep0, + noise_seed=args.common_best_noise_seed, + gpus=args.gpus, workers=args.workers, + ) + extension = _extend_common_winner_to_r100( + common, output=output / "selected_creative_full", + eval_ep0=args.full_eval_ep0, + confirm_ep0=args.full_confirm_ep0, + noise_seed=args.full_noise_seed, + gpus=args.gpus, workers=args.workers, + ) + achieved = bool( + extension is not None and extension["objective_achieved"] + ) + result = { + "status": STATUS, + "action": ( + "STOP_GOAL_ACHIEVED_BY_FULL_CREATIVE" + if achieved else "STOP_NO_VALIDATED_CREATIVE_ARM" + ), + "objective_achieved": achieved, + "stages_completed": ["B", "A", "C3"], + "stage_B": _ref(stage_b_root / "STAGE_COMPLETE.json"), + "stage_B_Dplus_only": ( + None if stage_b_causal is None else _ref( + output / "stage_B_Dplus_only" / "STAGE_COMPLETE.json" + ) + ), + "stage_A": _ref(stage_a_root / "STAGE_COMPLETE.json"), + "stage_C3": _ref(stage_c3_root / "STAGE_COMPLETE.json"), + "best_available_common_M50": _ref( + output / "common_best_M50" / "COMMON_BEST_COMPLETE.json" + ), + "best_available": common["selected"], + "full_creative_continuation": extension, + "completed_at": datetime.now(timezone.utc).isoformat(), + } + _write_final(output, result, source) + return result + + stage_c5_root = output / "stage_C_encoder_unfreeze_r5" + command = _training_command( + checkpoint=final_delivery["checkpoint"], output=stage_c5_root, + name=f"{selected_arm}_encoder01x", rounds=5, + scenario_ep0=scenario_ep0, eval_ep0=banks[3]["ep0"], + eval_rounds=(0, 3, 4, 5), lr=cfg["lr"], + inner_steps=cfg["inner_steps"], ell=cfg["ell"], + gp_cap=cfg["gp_cap"], selector="margin", + encoder_lr_ratio=0.1, workers=args.workers, + noise_seed=args.stage_c5_noise_seed, resume=stage_c3_root, + sample_seed=cfg["sample_seed"], audit_seed=cfg["audit_seed"], + train_seed=cfg["train_seed"], probe_seed=cfg["probe_seed"], + ) + _run( + command, log=output / "logs" / "stage_C5.log", + gpu=args.gpus[-1], cpu_range="96-159", + ) + c5_delivery = _read(stage_c5_root / "DELIVERY_COMPLETE.json") + c5_rows = _delivery_eval_rows(c5_delivery) + c5_controls = _evaluate_specs( + _base_specs(catalog, (3, 4, 5), "frozen_control"), + root=stage_c5_root / "matched_frozen_control", bank=banks[3], + noise_seed=args.stage_c5_noise_seed, gpus=[args.gpus[-1]], + workers=args.workers, + ) + r100_c5 = _evaluate_specs([{ + "name": "r100_reference", "round": 100, + "phase": "post_D0", "checkpoint": catalog[100]["post_D0"], + }], root=stage_c5_root / "r100_reference", bank=banks[3], + noise_seed=args.stage_c5_noise_seed, gpus=[args.gpus[-1]], + workers=args.workers)[0] + c5_increment_gate = _encoder_stage_gate( + c5_delivery, c5_rows, c5_controls, + {round_i: catalog[round_i]["record"] for round_i in (4, 5)}, + ) + combined_diagnostics = [ + *c3_gate["diagnostics"], *c5_increment_gate["diagnostics"], + ] + c5_gate = { + "rule": c5_increment_gate["rule"], + "diagnostics": combined_diagnostics, + "unsafe": bool(not any(row["eligible"] for row in combined_diagnostics)), + "signal": bool(any(row["signal"] for row in combined_diagnostics)), + "eligible_rounds": [ + row["round"] for row in combined_diagnostics if row["eligible"] + ], + } + c5_gate["continue_to_round5"] = bool( + not c5_gate["unsafe"] and c5_gate["signal"] + ) + selector_c5 = _selector_gate( + c5_rows, c5_controls, r100_c5, + diagnostic_rounds=c5_gate["eligible_rounds"], + ) + passed = bool(selector_c5["passed"]) + c5 = { + "status": STAGE_C_STATUS, + "phase": "rounds_4_to_5_after_authenticated_r3_resume", + "single_change": "enc_grid lr = 0.1 * trunk lr", + "training_delivery": _ref(stage_c5_root / "DELIVERY_COMPLETE.json"), + "bank": banks[3], + "rows": c5_rows, + "matched_frozen_controls": c5_controls, + "diagnostic_gate": c5_gate, + "promotion_gate": selector_c5, + "passed": passed, + "confirmation": None, + } + if passed: + c5["confirmation"] = _confirm_candidate( + training_run=stage_c5_root, + arm_name=c5_delivery["config"]["name"], + output=stage_c5_root / "confirmation", + ep0=args.stage_c_confirm_ep0, + noise_seed=args.stage_c_confirm_noise_seed, + gpus=args.gpus, workers=args.workers, + selected=selector_c5["selected"], + ) + creative_locks.append(_expanded_lock( + c5["confirmation"]["gamma_delivery"], + name="encoder_unfreeze_01x", + training_run=stage_c5_root, + )) + _write(stage_c5_root / "STAGE_COMPLETE.json", c5) + objective_achieved = bool( + c5["confirmation"] is not None + and c5["confirmation"]["objective_achieved"] + ) + common = None + extension = None + if not objective_achieved: + common = _common_best_available( + [baseline_lock, *creative_locks], + output=output / "common_best_M50", + ep0=args.common_best_ep0, + noise_seed=args.common_best_noise_seed, + gpus=args.gpus, workers=args.workers, + ) + extension = _extend_common_winner_to_r100( + common, output=output / "selected_creative_full", + eval_ep0=args.full_eval_ep0, + confirm_ep0=args.full_confirm_ep0, + noise_seed=args.full_noise_seed, + gpus=args.gpus, workers=args.workers, + ) + objective_achieved = bool( + extension is not None and extension["objective_achieved"] + ) + result = { + "status": STATUS, + "action": ( + "STOP_GOAL_ACHIEVED_BY_STAGE_C" + if objective_achieved else "STOP_NO_VALIDATED_CREATIVE_ARM" + ), + "stages_completed": ["B", "A", "C3", "C5"], + "stage_B": _ref(stage_b_root / "STAGE_COMPLETE.json"), + "stage_B_Dplus_only": ( + None if stage_b_causal is None else _ref( + output / "stage_B_Dplus_only" / "STAGE_COMPLETE.json" + ) + ), + "stage_A": _ref(stage_a_root / "STAGE_COMPLETE.json"), + "stage_C3": _ref(stage_c3_root / "STAGE_COMPLETE.json"), + "stage_C5": _ref(stage_c5_root / "STAGE_COMPLETE.json"), + "M10_qualified": bool(passed), + "objective_achieved": objective_achieved, + "full_creative_continuation": extension, + "best_available_common_M50": ( + None if common is None else _ref( + output / "common_best_M50" / "COMMON_BEST_COMPLETE.json" + ) + ), + "best_available": None if common is None else common["selected"], + "started_at": started.isoformat(), + "completed_at": datetime.now(timezone.utc).isoformat(), + } + _write_final(output, result, source) + return result + + +def build_parser(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--trigger", required=True) + parser.add_argument("--expected-trigger-source", required=True) + parser.add_argument("--output-dir", required=True) + parser.add_argument("--poll-seconds", type=int, default=60) + parser.add_argument("--gpus", type=int, nargs="+", default=(1, 3)) + parser.add_argument("--workers", type=int, default=32) + parser.add_argument( + "--audit-rounds", + default=",".join(map(str, DEFAULT_AUDIT_ROUNDS)), + ) + parser.add_argument("--stage-b-ep0", type=int, default=530_000) + parser.add_argument("--stage-a-ep0", type=int, default=540_000) + parser.add_argument("--stage-c3-ep0", type=int, default=550_000) + parser.add_argument("--stage-c5-ep0", type=int, default=560_000) + parser.add_argument("--stage-b-causal-ep0", type=int, default=535_000) + parser.add_argument("--stage-b-confirm-ep0", type=int, default=570_000) + parser.add_argument("--stage-a-confirm-ep0", type=int, default=580_000) + parser.add_argument("--stage-c-confirm-ep0", type=int, default=590_000) + parser.add_argument("--common-best-ep0", type=int, default=600_000) + parser.add_argument("--full-eval-ep0", type=int, default=610_000) + parser.add_argument("--full-confirm-ep0", type=int, default=620_000) + parser.add_argument("--stage-b-noise-seed", type=int, default=2_026_074_1) + parser.add_argument("--stage-a-noise-seed", type=int, default=2_026_074_2) + parser.add_argument("--stage-c3-noise-seed", type=int, default=2_026_074_3) + parser.add_argument("--stage-c5-noise-seed", type=int, default=2_026_074_4) + parser.add_argument( + "--stage-b-causal-noise-seed", type=int, default=2_026_074_5, + ) + parser.add_argument( + "--stage-b-confirm-noise-seed", type=int, default=2_026_074_6, + ) + parser.add_argument( + "--stage-a-confirm-noise-seed", type=int, default=2_026_074_7, + ) + parser.add_argument( + "--stage-c-confirm-noise-seed", type=int, default=2_026_074_8, + ) + parser.add_argument( + "--common-best-noise-seed", type=int, default=2_026_074_9, + ) + parser.add_argument("--full-noise-seed", type=int, default=2_026_075_0) + return parser + + +if __name__ == "__main__": + result = run(build_parser().parse_args()) + print(json.dumps(result, indent=2)) diff --git a/overnight_run_07_12_sfm/run_sfm_neutral_gamma_temperature.py b/overnight_run_07_12_sfm/run_sfm_neutral_gamma_temperature.py new file mode 100644 index 0000000..292be63 --- /dev/null +++ b/overnight_run_07_12_sfm/run_sfm_neutral_gamma_temperature.py @@ -0,0 +1,876 @@ +#!/usr/bin/env python3 +"""Calibrate a per-gamma raw temperature schedule, then reconfirm once. + +The preceding global-temperature M50 is explicitly reclassified as the +calibration bank. A seven-value schedule is selected independently for the +pretrained and expanded policies, hashed, and then evaluated on one fresh +disjoint M50 bank. Locked Kazuki is never retuned. +""" +from __future__ import annotations + +import argparse +from datetime import datetime, timezone +import hashlib +import itertools +import json +from pathlib import Path +import subprocess +import sys +import time + +import numpy as np + +import run_sfm_neutral_temperature_m50 as BASE +import sfm_protocol as SP + + +STATUS = "SFM_NEUTRAL_GAMMA_TEMPERATURE_M50_COMPLETE" + + +def _gpu_inventory(indices: list[int]) -> dict: + output = subprocess.check_output([ + "nvidia-smi", + "--query-gpu=index,uuid,name,driver_version,memory.total", + "--format=csv,noheader,nounits", + ], text=True) + rows = [] + for line in output.splitlines(): + fields = [field.strip() for field in line.split(",")] + if len(fields) != 5: + raise RuntimeError("unexpected nvidia-smi inventory row") + rows.append({ + "index": int(fields[0]), "uuid": fields[1], "name": fields[2], + "driver_version": fields[3], "memory_total_MiB": int(fields[4]), + }) + by_index = {row["index"]: row for row in rows} + if any(int(index) not in by_index for index in indices): + raise RuntimeError("requested physical GPU is absent") + return {"requested_indices": list(map(int, indices)), "devices": rows} + + +def _wait(path: Path, poll: int) -> dict: + while not path.is_file(): + print(f"WAITING {path}", flush=True) + time.sleep(int(poll)) + return BASE._read(path) + + +def _point(rows: list[dict]) -> dict: + if not rows: + raise ValueError("empty evaluation rows") + successes = [row for row in rows if row["success"]] + return { + "SR": float(np.mean([row["success"] for row in rows])), + "CR": float(np.mean([row["collision"] for row in rows])), + "timeout": float(np.mean([row["timeout"] for row in rows])), + "Validity": float(np.mean([row["validity"] for row in rows])), + "clearance": ( + float(np.mean([row["successful_clearance"] for row in successes])) + if successes else float("nan") + ), + "time_to_goal": ( + float(np.mean([row["time_to_goal"] for row in successes])) + if successes else float("nan") + ), + } + + +def _candidate(method: str, schedule: tuple[float, ...], payloads: dict[float, dict], + checkpoint: str, round_index: int) -> dict: + rows, per_gamma = [], {} + for gamma, temperature in zip(SP.GAMMAS, schedule): + payload = payloads[float(temperature)] + source = payload["records"][0]["cell"]["rows"] + selected = [row for row in source if float(row["gamma"]) == float(gamma)] + if len(selected) != 50: + raise RuntimeError("calibration cell lacks M50/gamma rows") + rows.extend(selected) + per_gamma[str(gamma)] = _point(selected) + return { + "method": method, + "round": int(round_index), + "checkpoint": checkpoint, + "temperature": None, + "temperature_by_gamma": list(schedule), + "pooled": _point(rows), + "per_gamma": per_gamma, + } + + +def _enumerate(method: str, payloads: dict[float, dict], temperatures: list[float], + checkpoint: str, round_index: int): + for schedule in itertools.product(temperatures, repeat=len(SP.GAMMAS)): + yield _candidate(method, schedule, payloads, checkpoint, round_index) + + +def _single_reference_target(reference: dict) -> dict: + return {metric: reference["pooled"][metric] for metric in BASE.METRICS} + + +def _pick(candidates, *, target: dict, liveness: dict) -> dict: + def key(row): + trend = BASE._trend(row) + shortfall = BASE._shortfalls(row, target) + return ( + 0 if BASE._liveness_eligible(row, liveness) else 1, + 0 if BASE._trend_eligible(row) else 1, + max(shortfall.values()), + sum(shortfall.values()), + -trend["mean_fraction"], + -row["pooled"]["SR"], + row["pooled"]["timeout"], + tuple(row["temperature_by_gamma"]), + ) + return min( + candidates, + key=key, + ) + + +def _schedule_sha(record: dict) -> str: + value = json.dumps({ + "checkpoint_sha256": record["checkpoint_sha256"], + "round": record["round"], + "temperature_by_gamma": record["temperature_by_gamma"], + "gammas": list(map(float, SP.GAMMAS)), + }, sort_keys=True, separators=(",", ":")).encode() + return hashlib.sha256(value).hexdigest() + + +def _final_objective_gates(final_records: list[dict], comparisons: dict) -> dict: + by_method = {row["method"]: row for row in final_records} + pretrained = by_method["pretrained"] + expanded = by_method["expanded"] + kazuki = by_method["kazuki_locked"] + liveness_contract = BASE._liveness_contract([pretrained], kazuki) + expanded_trend = BASE._trend(expanded) + paired_ci_clean = all(map(_ci_win, comparisons.values())) + liveness_eligible = BASE._liveness_eligible(expanded, liveness_contract) + gamma_trend_eligible = BASE._trend_eligible(expanded) + return { + "paired_ci_clean_four_metric_win": paired_ci_clean, + "final_liveness_contract": liveness_contract, + "final_liveness_eligible": liveness_eligible, + "final_gamma_trend_eligible": gamma_trend_eligible, + "objective_achieved": ( + paired_ci_clean and liveness_eligible and gamma_trend_eligible + ), + } + + +def _raw_rows(payload: dict) -> list[dict]: + return payload["records"][0]["cell"]["rows"] + + +def _rows_sha256(payload: dict) -> str: + encoded = json.dumps( + _raw_rows(payload), + sort_keys=True, + separators=(",", ":"), + allow_nan=False, + ).encode() + return hashlib.sha256(encoded).hexdigest() + + +def _json_value_sha256(value) -> str: + encoded = json.dumps( + value, sort_keys=True, separators=(",", ":"), allow_nan=False, + ).encode() + return hashlib.sha256(encoded).hexdigest() + + +def _reuse_global_temperature_cells( + initial_path: Path, + initial: dict, + methods: dict[str, dict], +) -> tuple[dict[str, dict[float, dict]], dict]: + bank = initial["banks"]["disjoint_confirmation"] + cells = {method: {} for method in methods} + records = {} + frozen_records = { + row["method"]: row for row in initial.get("final_records", ()) + } + for method in ("pretrained", "expanded"): + selected = methods[method] + temperature = float(selected["temperature"]) + reference_path = ( + initial_path.parent / "disjoint_m50" / method + / "raw_m50_offline_metrics.json" + ) + payload = BASE._read(reference_path) + record = payload["records"][0] + rows = _raw_rows(payload) + expected_ids = list(range( + int(bank["ep0"]), int(bank["ep0"]) + int(bank["M_per_gamma"]) + )) + noise = payload.get("noise_bank", {}) + per_gamma_ids = { + float(gamma): sorted( + int(row["episode"]) for row in rows + if float(row["gamma"]) == float(gamma) + ) + for gamma in SP.GAMMAS + } + observed_schedule = ( + payload.get("temperature_by_gamma") + or noise.get("temperature_by_gamma") + or [float(payload["temperature"])] * len(SP.GAMMAS) + ) + if ( + payload.get("scene_profile") != "double_density_velocity_ood" + or record["cell"].get("scene_profile") + != "double_density_velocity_ood" + or int(payload["bank"]["ep0"]) != int(bank["ep0"]) + or int(payload["bank"]["M_per_gamma"]) != int(bank["M_per_gamma"]) + or payload["bank"].get("scenario_ids") != expected_ids + or int(payload["noise_bank"]["seed"]) != int(bank["noise_seed"]) + or noise.get("gammas") != list(map(float, SP.GAMMAS)) + or int(noise.get("NFE", -1)) != 8 + or noise.get("dtype") != "float32" + or noise.get("shape") != [ + len(SP.GAMMAS), int(bank["M_per_gamma"]), 180, 20 + ] + or not isinstance(noise.get("sha256"), str) + or len(noise["sha256"]) != 64 + or float(payload["temperature"]) != temperature + or observed_schedule != [temperature] * len(SP.GAMMAS) + or int(record["round"]) != int(selected["round"]) + or record["cell"]["checkpoint_sha256"] + != selected["checkpoint_sha256"] + or len(rows) != int(bank["M_per_gamma"]) * len(SP.GAMMAS) + or any(ids != expected_ids for ids in per_gamma_ids.values()) + ): + raise RuntimeError( + f"global-temperature reference contract failed for {method}" + ) + observed = BASE._record_from_cell( + payload, method=method, temperature=temperature, + ) + frozen = frozen_records.get(method) + if frozen is None or any( + observed[key] != frozen[key] + for key in ( + "round", "checkpoint", "checkpoint_sha256", "temperature", + "pooled", "per_gamma", + ) + ): + raise RuntimeError( + f"global-temperature sidecar changed after delivery for {method}" + ) + cells[method][temperature] = payload + records[method] = { + "temperature": temperature, + "reference": str(reference_path), + "reference_file_sha256": BASE.FUNNEL.sha256_file(reference_path), + "reference_rows_sha256": _rows_sha256(payload), + "noise_bank_sha256": noise["sha256"], + "cell_key": record["cell"]["cell_key"], + "checkpoint_sha256": selected["checkpoint_sha256"], + "rerun": False, + } + if ( + records["pretrained"]["noise_bank_sha256"] + != records["expanded"]["noise_bank_sha256"] + ): + raise RuntimeError( + "reused pretrained and expanded cells do not share CRN bytes" + ) + reuse = { + "status": "GLOBAL_TEMPERATURE_CALIBRATION_CELLS_REUSED", + "reason": ( + "the locked global-temperature M50 cells are already the exact " + "calibration-bank observations; re-running a GPU rollout can " + "perturb borderline trajectories and is neither required nor " + "scientifically preferable" + ), + "records": records, + } + return cells, reuse + + +def _validate_prior_calibration_cell( + path: Path, payload: dict, *, method: str, temperature: float, + selected: dict, bank: dict, +) -> dict: + record = payload["records"][0] + rows = _raw_rows(payload) + expected_ids = list(range( + int(bank["ep0"]), int(bank["ep0"]) + int(bank["M_per_gamma"]) + )) + noise = payload.get("noise_bank", {}) + per_gamma_ids = { + float(gamma): sorted( + int(row["episode"]) for row in rows + if float(row["gamma"]) == float(gamma) + ) + for gamma in SP.GAMMAS + } + if ( + payload.get("scene_profile") != "double_density_velocity_ood" + or record["cell"].get("scene_profile") + != "double_density_velocity_ood" + or int(payload["bank"]["ep0"]) != int(bank["ep0"]) + or int(payload["bank"]["M_per_gamma"]) != int(bank["M_per_gamma"]) + or payload["bank"].get("scenario_ids") != expected_ids + or int(noise.get("seed", -1)) != int(bank["noise_seed"]) + or noise.get("gammas") != list(map(float, SP.GAMMAS)) + or int(noise.get("NFE", -1)) != 8 + or noise.get("dtype") != "float32" + or noise.get("shape") != [ + len(SP.GAMMAS), int(bank["M_per_gamma"]), 180, 20 + ] + or not isinstance(noise.get("sha256"), str) + or len(noise["sha256"]) != 64 + or float(payload["temperature"]) != float(temperature) + or payload.get("temperature_by_gamma") + != [float(temperature)] * len(SP.GAMMAS) + or int(record["round"]) != int(selected["round"]) + or record["cell"]["checkpoint_sha256"] + != selected["checkpoint_sha256"] + or len(rows) != int(bank["M_per_gamma"]) * len(SP.GAMMAS) + or any(ids != expected_ids for ids in per_gamma_ids.values()) + ): + raise RuntimeError( + f"prior calibration contract failed for {method} temp={temperature:g}" + ) + return { + "method": method, + "temperature": float(temperature), + "reference": str(path), + "reference_file_sha256": BASE.FUNNEL.sha256_file(path), + "reference_rows_sha256": _rows_sha256(payload), + "noise_bank_sha256": noise["sha256"], + "cell_key": record["cell"]["cell_key"], + "checkpoint_sha256": selected["checkpoint_sha256"], + } + + +def _reuse_prior_calibration_cells( + root: Path, methods: dict[str, dict], bank: dict, + temperatures: list[float], cells: dict[str, dict[float, dict]], +) -> dict: + records = [] + for method, selected in methods.items(): + for temperature in temperatures: + if float(temperature) in cells[method]: + continue + name = f"{method}_temp{temperature:g}".replace(".", "p") + path = root / "calibration_m50" / name / "raw_m50_offline_metrics.json" + if not path.is_file(): + continue + payload, file_sha = BASE._read_hashed(path) + provenance = _validate_prior_calibration_cell( + path, payload, method=method, temperature=temperature, + selected=selected, bank=bank, + ) + if provenance["reference_file_sha256"] != file_sha: + raise RuntimeError("calibration bytes changed while being read") + cells[method][float(temperature)] = payload + records.append(provenance) + + for temperature in temperatures: + if all(float(temperature) in cells[method] for method in methods): + hashes = { + cells[method][float(temperature)]["noise_bank"]["sha256"] + for method in methods + } + if len(hashes) != 1: + raise RuntimeError( + f"reused temp={temperature:g} cells do not share CRN bytes" + ) + return { + "status": "PRIOR_CALIBRATION_CELLS_CONTENT_AUTHENTICATED", + "root": str(root), + "records": records, + "missing_cells_will_be_run": [ + {"method": method, "temperature": float(temperature)} + for method in methods for temperature in temperatures + if float(temperature) not in cells[method] + ], + } + + +def _calibration_crn_sha( + cells: dict[str, dict[float, dict]], temperatures: list[float], + methods: dict[str, dict], +) -> str: + missing = [ + (method, float(temperature)) + for method in ("pretrained", "expanded") + for temperature in temperatures + if float(temperature) not in cells[method] + ] + if missing: + raise RuntimeError(f"calibration grid is incomplete: {missing}") + hashes = { + cells[method][float(temperature)]["noise_bank"]["sha256"] + for method in ("pretrained", "expanded") + for temperature in temperatures + } + if len(hashes) != 1: + raise RuntimeError("calibration temperatures do not share one CRN bank") + for method, selected in methods.items(): + for temperature in temperatures: + observed = cells[method][float(temperature)]["records"][0][ + "cell" + ]["checkpoint_sha256"] + if observed != selected["checkpoint_sha256"]: + raise RuntimeError("calibration cell checkpoint digest changed") + return next(iter(hashes)) + + +def _authenticate_method_checkpoints(methods: dict[str, dict]) -> None: + for method, selected in methods.items(): + path = Path(selected["checkpoint"]).resolve() + if BASE.FUNNEL.sha256_file(path) != selected["checkpoint_sha256"]: + raise RuntimeError(f"locked {method} checkpoint digest mismatch") + + +def _validate_fresh_raw_cell( + payload: dict, *, locked: dict, ep0: int, noise_seed: int, +) -> str: + record = payload["records"][0] + rows = _raw_rows(payload) + expected_ids = list(range(int(ep0), int(ep0) + 50)) + noise = payload.get("noise_bank", {}) + per_gamma_ids = { + float(gamma): sorted( + int(row["episode"]) for row in rows + if float(row["gamma"]) == float(gamma) + ) + for gamma in SP.GAMMAS + } + if ( + payload.get("scene_profile") != "double_density_velocity_ood" + or int(payload["bank"]["ep0"]) != int(ep0) + or int(payload["bank"]["M_per_gamma"]) != 50 + or payload["bank"].get("scenario_ids") != expected_ids + or int(noise.get("seed", -1)) != int(noise_seed) + or noise.get("gammas") != list(map(float, SP.GAMMAS)) + or int(noise.get("NFE", -1)) != 8 + or noise.get("dtype") != "float32" + or noise.get("shape") != [len(SP.GAMMAS), 50, 180, 20] + or not isinstance(noise.get("sha256"), str) + or len(noise["sha256"]) != 64 + or payload.get("temperature_by_gamma") + != list(map(float, locked["temperature_by_gamma"])) + or int(record["round"]) != int(locked["round"]) + or record["cell"]["checkpoint_sha256"] + != locked["checkpoint_sha256"] + or len(rows) != 50 * len(SP.GAMMAS) + or any(ids != expected_ids for ids in per_gamma_ids.values()) + ): + raise RuntimeError("fresh raw confirmation contract changed") + return noise["sha256"] + + +def _validate_fresh_kazuki( + payload: dict, *, pretrained: dict, ep0: int, +) -> None: + expected_ids = list(range(int(ep0), int(ep0) + 50)) + rows = payload.get("rows", ()) + per_gamma_ids = { + float(gamma): sorted( + int(row["episode"]) for row in rows + if float(row["gamma"]) == float(gamma) + ) + for gamma in SP.GAMMAS + } + if ( + payload.get("scene_profile") != "double_density_velocity_ood" + or int(payload.get("ep0", -1)) != int(ep0) + or int(payload.get("M_per_gamma", -1)) != 50 + or payload.get("checkpoint_sha256") + != pretrained["checkpoint_sha256"] + or len(rows) != 50 * len(SP.GAMMAS) + or any(ids != expected_ids for ids in per_gamma_ids.values()) + ): + raise RuntimeError("fresh Kazuki confirmation contract changed") + + +def _validate_calibration_kazuki( + payload: dict, *, initial: dict, bank: dict, pretrained: dict, +) -> dict: + frozen = next( + (row for row in initial.get("final_records", ()) + if row.get("method") == "kazuki_locked"), + None, + ) + observed = BASE._kazuki_record(payload) + if frozen is None or any( + observed[key] != frozen[key] + for key in ( + "round", "checkpoint", "checkpoint_sha256", "temperature", + "pooled", "per_gamma", + ) + ): + raise RuntimeError("calibration Kazuki sidecar changed after delivery") + _validate_fresh_kazuki( + payload, pretrained=pretrained, ep0=int(bank["ep0"]), + ) + if int(bank["M_per_gamma"]) != 50: + raise RuntimeError("Kazuki calibration bank is not M50") + return observed + + +def _metric(rows: list[dict], name: str) -> float: + point = _point(rows) + return float(point[name]) + + +def _paired_cluster_ci(expanded: list[dict], reference: list[dict], *, seed: int, + draws: int = 5_000) -> dict: + exp_by_episode, ref_by_episode = {}, {} + for row in expanded: + exp_by_episode.setdefault(int(row["episode"]), []).append(row) + for row in reference: + ref_by_episode.setdefault(int(row["episode"]), []).append(row) + episodes = sorted(set(exp_by_episode) & set(ref_by_episode)) + if len(episodes) != 50 or any( + len(exp_by_episode[key]) != 7 or len(ref_by_episode[key]) != 7 + for key in episodes + ): + raise RuntimeError("paired CI requires 50 complete scenario clusters") + generator = np.random.default_rng(int(seed)) + samples = {metric: [] for metric in BASE.METRICS} + for _ in range(int(draws)): + chosen = generator.choice(episodes, size=len(episodes), replace=True) + exp_rows = [row for key in chosen for row in exp_by_episode[int(key)]] + ref_rows = [row for key in chosen for row in ref_by_episode[int(key)]] + for metric in BASE.METRICS: + samples[metric].append( + _metric(exp_rows, metric) - _metric(ref_rows, metric) + ) + result = {} + for metric, values in samples.items(): + values = np.asarray(values, float) + finite = values[np.isfinite(values)] + if not len(finite): + result[metric] = {"mean_delta": None, "paired_cluster_95": [None, None]} + else: + result[metric] = { + "mean_delta": float(np.mean(finite)), + "paired_cluster_95": list(map( + float, np.quantile(finite, [.025, .975]) + )), + "desired_sign": "negative" if metric in BASE.LOWER_IS_BETTER else "positive", + } + return result + + +def _ci_win(value: dict) -> bool: + for metric, row in value.items(): + low, high = row["paired_cluster_95"] + if low is None: + return False + if metric in BASE.LOWER_IS_BETTER: + if not high < 0.0: + return False + elif not low > 0.0: + return False + return True + + +def run(args) -> dict: + started = datetime.now(timezone.utc) + source_commit = BASE._source_gate(args.expected_source_commit) + gpu_provenance = _gpu_inventory(args.gpus) + initial_path = Path(args.initial_delivery).resolve() + _wait(initial_path, args.poll_seconds) + initial, initial_sha256 = BASE._read_hashed(initial_path) + if initial.get("status") != BASE.STATUS: + raise RuntimeError("invalid global-temperature delivery") + calibration_bank = initial["banks"]["disjoint_confirmation"] + if int(calibration_bank["M_per_gamma"]) != 50: + raise RuntimeError("initial result is not M50") + final_ids = set(range(int(args.final_ep0), int(args.final_ep0) + 50)) + calibration_ids = set(range( + int(calibration_bank["ep0"]), int(calibration_bank["ep0"]) + 50 + )) + if final_ids & calibration_ids: + raise RuntimeError("fresh confirmation overlaps calibration M50") + + output = Path(args.output_dir).resolve() + if output.exists(): + raise FileExistsError(output) + output.mkdir(parents=True) + temperatures = list(map(float, args.temperatures.split(","))) + if temperatures != [0.55, 0.7, 0.85, 1.0]: + raise ValueError("calibration grid changed") + selected = initial["selection"] + methods = { + "pretrained": selected["selected_pretrained"], + "expanded": selected["selected_expanded"], + } + _authenticate_method_checkpoints(methods) + cells, reuse = _reuse_global_temperature_cells( + initial_path, initial, methods + ) + BASE._write(output / "GLOBAL_TEMPERATURE_REUSE.json", reuse) + prior_reuse = { + "status": "PRIOR_CALIBRATION_REUSE_NOT_REQUESTED", "records": [], + } + if args.reuse_calibration_root: + prior_reuse = _reuse_prior_calibration_cells( + Path(args.reuse_calibration_root).resolve(), methods, + calibration_bank, temperatures, cells, + ) + BASE._write(output / "PRIOR_CALIBRATION_REUSE.json", prior_reuse) + + calibration_root = output / "calibration_m50" + jobs, metadata = [], {} + for method, record in methods.items(): + for temperature in temperatures: + if temperature in cells[method]: + continue + name = f"{method}_temp{temperature:g}".replace(".", "p") + out = calibration_root / name + jobs.append({ + "name": name, + "command": BASE._raw_command( + record["checkpoint"], round_index=record["round"], + temperature=temperature, + ep0=int(calibration_bank["ep0"]), + noise_seed=int(calibration_bank["noise_seed"]), M=50, + workers=args.workers, output=out, + cache=calibration_root / "cache", + ), + }) + metadata[name] = (method, temperature, out) + BASE._run_jobs( + jobs, gpus=args.gpus, workers=args.workers, log_dir=output / "logs" + ) + for method, temperature, out in metadata.values(): + cells[method][temperature] = BASE._read( + out / "raw_m50_offline_metrics.json" + ) + calibration_noise_sha256 = _calibration_crn_sha( + cells, temperatures, methods, + ) + _authenticate_method_checkpoints(methods) + + kazuki_path = initial_path.parent / "disjoint_m50" / "kazuki_locked.json" + kazuki_payload = BASE._read(kazuki_path) + kazuki = _validate_calibration_kazuki( + kazuki_payload, initial=initial, bank=calibration_bank, + pretrained=methods["pretrained"], + ) + kazuki_calibration_reference = { + "path": str(kazuki_path), + "file_sha256": BASE.FUNNEL.sha256_file(kazuki_path), + "rows_sha256": _json_value_sha256(kazuki_payload["rows"]), + } + kazuki_liveness = BASE._liveness_contract([], kazuki) + pretrained = _pick( + _enumerate( + "pretrained", cells["pretrained"], temperatures, + methods["pretrained"]["checkpoint"], methods["pretrained"]["round"], + ), + target=_single_reference_target(kazuki), + liveness=kazuki_liveness, + ) + pretrained["checkpoint_sha256"] = methods["pretrained"]["checkpoint_sha256"] + target = BASE._envelope([pretrained], kazuki) + liveness = BASE._liveness_contract([pretrained], kazuki) + expanded = _pick( + _enumerate( + "expanded", cells["expanded"], temperatures, + methods["expanded"]["checkpoint"], methods["expanded"]["round"], + ), + target=target, + liveness=liveness, + ) + expanded["checkpoint_sha256"] = methods["expanded"]["checkpoint_sha256"] + lock = { + "status": "GAMMA_TEMPERATURE_SCHEDULE_LOCKED", + "calibration_bank_role": ( + "the former global-temperature M50 is intentionally reused and " + "therefore reclassified as calibration, not confirmation" + ), + "analysis_source_commit": source_commit, + "initial_delivery": str(initial_path), + "initial_delivery_sha256": initial_sha256, + "temperature_grid": temperatures, + "calibration_noise_bank_sha256": calibration_noise_sha256, + "pretrained": pretrained, + "expanded": expanded, + "kazuki": kazuki, + "kazuki_calibration_reference": kazuki_calibration_reference, + "target_envelope": target, + "liveness_contract": liveness, + "gamma_trend_gate": { + "minimum_each_adjacent_pair_family_fraction": ( + BASE.MIN_TREND_FAMILY_FRACTION + ), + "pretrained": BASE._trend(pretrained), + "expanded": BASE._trend(expanded), + }, + "expanded_shortfalls": BASE._shortfalls(expanded, target), + "schedule_sha256": { + "pretrained": _schedule_sha(pretrained), + "expanded": _schedule_sha(expanded), + }, + } + BASE._write(output / "SCHEDULE_LOCKED.json", lock) + + final_root = output / "fresh_disjoint_m50" + _authenticate_method_checkpoints(methods) + final_jobs, final_meta = [], {} + for method, record in (("pretrained", pretrained), ("expanded", expanded)): + out = final_root / method + final_jobs.append({ + "name": f"fresh_{method}", + "command": BASE.FUNNEL.evaluator_command( + [record["checkpoint"]], [f"r{record['round']}"], + scene_profile="double_density_velocity_ood", + ep0=args.final_ep0, noise_seed=args.final_noise_seed, + m_per_gamma=50, workers=args.workers, + cache_dir=final_root / "cache", output_dir=out, + temperature=1.0, + temperature_by_gamma=record["temperature_by_gamma"], + ), + }) + final_meta[method] = out + final_kazuki = final_root / "kazuki_locked.json" + final_jobs.append({ + "name": "fresh_kazuki", + "command": [ + sys.executable, str(Path(BASE.__file__).resolve()), "kazuki", + "--checkpoint", methods["pretrained"]["checkpoint"], + "--ep0", str(args.final_ep0), "--M", "50", + "--workers", str(args.workers), "--device", "cuda:0", + "--output", str(final_kazuki), + "--expected-source-commit", source_commit, + ], + }) + confirmation_gpu = int(args.gpus[0]) + BASE._run_jobs( + final_jobs, gpus=[confirmation_gpu], workers=args.workers, + log_dir=output / "logs", + ) + raw_payloads = { + method: BASE._read(path / "raw_m50_offline_metrics.json") + for method, path in final_meta.items() + } + final_kazuki_payload = BASE._read(final_kazuki) + _authenticate_method_checkpoints(methods) + fresh_noise_hashes = { + _validate_fresh_raw_cell( + raw_payloads[method], locked=locked, + ep0=args.final_ep0, noise_seed=args.final_noise_seed, + ) + for method, locked in (("pretrained", pretrained), ("expanded", expanded)) + } + if len(fresh_noise_hashes) != 1: + raise RuntimeError("fresh raw methods do not share one CRN bank") + _validate_fresh_kazuki( + final_kazuki_payload, pretrained=methods["pretrained"], + ep0=args.final_ep0, + ) + final_records = [ + BASE._record_from_cell( + raw_payloads[method], method=method, temperature=1.0 + ) + for method in ("pretrained", "expanded") + ] + for row, locked in zip(final_records, (pretrained, expanded)): + row["temperature"] = None + row["temperature_by_gamma"] = locked["temperature_by_gamma"] + final_records.append(BASE._kazuki_record(final_kazuki_payload)) + plot_path = final_root / "four_metric_per_gamma.png" + BASE._render(final_records, plot_path) + exp_rows = _raw_rows(raw_payloads["expanded"]) + comparisons = { + "expanded_minus_pretrained": _paired_cluster_ci( + exp_rows, _raw_rows(raw_payloads["pretrained"]), + seed=args.final_noise_seed + 1, + ), + "expanded_minus_kazuki": _paired_cluster_ci( + exp_rows, final_kazuki_payload["rows"], + seed=args.final_noise_seed + 2, + ), + } + objective_gates = _final_objective_gates(final_records, comparisons) + if BASE.FUNNEL.sha256_file(initial_path) != initial_sha256: + raise RuntimeError("initial delivery changed during gamma calibration") + BASE._source_gate(source_commit) + fresh_artifacts = {} + for method, path in final_meta.items(): + metrics_path = path / "raw_m50_offline_metrics.json" + payload = raw_payloads[method] + fresh_artifacts[method] = { + "path": str(metrics_path), + "file_sha256": BASE.FUNNEL.sha256_file(metrics_path), + "rows_sha256": _rows_sha256(payload), + "noise_bank_sha256": payload["noise_bank"]["sha256"], + "checkpoint_sha256": payload["records"][0]["cell"][ + "checkpoint_sha256" + ], + } + fresh_artifacts["kazuki_locked"] = { + "path": str(final_kazuki), + "file_sha256": BASE.FUNNEL.sha256_file(final_kazuki), + "rows_sha256": _json_value_sha256(final_kazuki_payload["rows"]), + "checkpoint_sha256": final_kazuki_payload["checkpoint_sha256"], + } + fresh_artifacts["plots"] = { + str(path): BASE.FUNNEL.sha256_file(path) + for path in (plot_path, plot_path.with_suffix(".pdf")) + } + result = { + "status": STATUS, + "source_commit": source_commit, + "initial_delivery": str(initial_path), + "initial_delivery_sha256": initial_sha256, + "global_temperature_calibration_reuse": reuse, + "prior_calibration_reuse": prior_reuse, + "gpu_provenance": gpu_provenance, + "calibration_bank": calibration_bank, + "fresh_confirmation_bank": { + "M_per_gamma": 50, + "ep0": args.final_ep0, + "noise_seed": args.final_noise_seed, + "single_physical_gpu_index": confirmation_gpu, + "execution": "sequential to avoid cross-GPU comparison noise", + }, + "fresh_confirmation_artifacts": fresh_artifacts, + "lock": lock, + "final_records": final_records, + "paired_cluster_differences": comparisons, + "ci_clean_four_metric_win": objective_gates[ + "paired_ci_clean_four_metric_win" + ], + **objective_gates, + "gamma_trends": { + row["method"]: BASE._trend(row) for row in final_records + }, + "completed_at": datetime.now(timezone.utc).isoformat(), + "wall_seconds": (datetime.now(timezone.utc) - started).total_seconds(), + } + BASE._write(output / "DELIVERY_COMPLETE.json", result) + print(json.dumps({ + "status": STATUS, + "ci_clean_four_metric_win": result["ci_clean_four_metric_win"], + "objective_achieved": result["objective_achieved"], + "delivery": str(output / "DELIVERY_COMPLETE.json"), + }, indent=2)) + return result + + +def build_parser(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--initial-delivery", required=True) + parser.add_argument("--output-dir", required=True) + parser.add_argument("--temperatures", default="0.55,0.7,0.85,1.0") + parser.add_argument("--gpus", nargs="+", type=int, default=[1, 3]) + parser.add_argument("--workers", type=int, default=32) + parser.add_argument("--poll-seconds", type=int, default=60) + parser.add_argument("--final-ep0", type=int, default=480_000) + parser.add_argument("--final-noise-seed", type=int, default=2_026_073_6) + parser.add_argument("--reuse-calibration-root") + parser.add_argument("--expected-source-commit") + return parser + + +if __name__ == "__main__": + run(build_parser().parse_args()) diff --git a/overnight_run_07_12_sfm/run_sfm_neutral_temperature_m50.py b/overnight_run_07_12_sfm/run_sfm_neutral_temperature_m50.py new file mode 100644 index 0000000..735ee41 --- /dev/null +++ b/overnight_run_07_12_sfm/run_sfm_neutral_temperature_m50.py @@ -0,0 +1,943 @@ +#!/usr/bin/env python3 +"""Validation-select one global raw temperature, then run disjoint M50. + +This is a post-training evaluator. It never mutates the four neutral-teacher +training runs. Temperature and checkpoint selection use only the completed +runs' M20 screen plus a fresh M10 validation bank. The final M50 scenario and +noise bank are not read until one expanded checkpoint/temperature and one +pretrained temperature are locked. +""" +from __future__ import annotations + +import argparse +from concurrent.futures import ProcessPoolExecutor, ThreadPoolExecutor +from datetime import datetime, timezone +import hashlib +import json +import math +import os +from pathlib import Path +import subprocess +import sys +import time + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np + +import grid_policy_sfm as GPS +import run_sfm_b1_offline_eval_funnel as FUNNEL +import sfm_b1_offline_eval as EVAL +import sfm_kazuki as KZ +import sfm_protocol as SP +import sfm_scene as SS + + +HERE = Path(__file__).resolve().parent +STATUS = "SFM_NEUTRAL_TEMPERATURE_DISJOINT_M50_COMPLETE" +METRICS = ("CR", "Validity", "clearance", "time_to_goal") +LOWER_IS_BETTER = {"CR", "time_to_goal"} +MIN_TREND_FAMILY_FRACTION = .75 + + +def _source_gate(expected: str | None = None) -> str: + commit = subprocess.check_output( + ["git", "rev-parse", "HEAD"], cwd=HERE, text=True, + ).strip() + dirty = subprocess.check_output( + ["git", "status", "--porcelain"], cwd=HERE, text=True, + ).strip() + if dirty: + raise RuntimeError("evaluation requires a clean frozen worktree") + if expected is not None and commit != str(expected): + raise RuntimeError("evaluation source commit differs from the pin") + return commit + + +def _read(path: Path) -> dict: + with path.open() as stream: + return json.load(stream) + + +def _read_hashed(path: Path) -> tuple[dict, str]: + """Read immutable JSON bytes once so payload and digest cannot diverge.""" + payload = path.read_bytes() + return json.loads(payload), hashlib.sha256(payload).hexdigest() + + +def _write(path: Path, value) -> None: + FUNNEL.write_json(path, value) + + +def _wait_for_deliveries(root: Path, arm_names: list[str], poll: int) -> list[dict]: + paths = [root / name / "DELIVERY_COMPLETE.json" for name in arm_names] + while True: + missing = [path for path in paths if not path.is_file()] + if not missing: + break + print(f"WAITING deliveries {len(paths) - len(missing)}/{len(paths)}", flush=True) + time.sleep(int(poll)) + payloads = [_read(path) for path in paths] + for path, payload in zip(paths, payloads): + if payload.get("status") != "SFM_B1_NEUTRAL_MULTIROUND_COMPLETE": + raise RuntimeError(f"invalid delivery: {path}") + checkpoints = {payload["checkpoint_sha256"] for payload in payloads} + if len(checkpoints) != 1: + raise RuntimeError("arms do not share one pretrained checkpoint") + return payloads + + +def _ref_identity(ref: dict) -> tuple: + return ( + str(Path(ref["path"]).resolve()), + ref["sha256"], + int(ref["round"]) if "round" in ref else None, + ) + + +def _round_ref_from_marker(path: Path) -> dict: + path = path.resolve() + record, marker_sha = _read_hashed(path) + if record.get("status") != "SFM_B1_NEUTRAL_MULTIROUND_ROUND_COMPLETE": + raise RuntimeError("round marker status changed") + checkpoint = Path(record["checkpoint"]).resolve() + checkpoint_sha = record["checkpoint_sha256"] + if FUNNEL.sha256_file(checkpoint) != checkpoint_sha: + raise RuntimeError("round checkpoint digest mismatch") + ref = { + "path": str(path), "sha256": marker_sha, + "round": int(record["round"]), + "post_D0": str(checkpoint), "post_D0_sha256": checkpoint_sha, + } + post_positive = checkpoint.with_name( + f"round_{int(record['round']):02d}_post_positive.pt" + ) + if post_positive.is_file(): + ref.update({ + "post_Dplus": str(post_positive.resolve()), + "post_Dplus_sha256": FUNNEL.sha256_file(post_positive), + }) + return ref + + +def _delivery_round_refs(delivery: dict, *, seen=()) -> list[dict]: + """Authenticate a possibly nested resume chain and return all refs.""" + current_paths = [ + str(Path(path).resolve()) for path in delivery.get("round_records", ()) + ] + current_refs = list(delivery.get("round_record_refs", ())) + if current_refs: + if [str(Path(ref["path"]).resolve()) for ref in current_refs] != current_paths: + raise RuntimeError("current round-record path snapshot changed") + else: + current_refs = [ + _round_ref_from_marker(Path(path)) for path in current_paths + ] + + prior_refs = [] + resume = delivery.get("resume") + if resume: + prior_path = Path(resume["delivery"]).resolve() + if str(prior_path) in seen: + raise RuntimeError("cyclic resume delivery chain") + prior, prior_sha = _read_hashed(prior_path) + if prior_sha != resume["delivery_sha256"]: + raise RuntimeError("resume delivery digest mismatch") + if prior.get("status") != "SFM_B1_NEUTRAL_MULTIROUND_COMPLETE": + raise RuntimeError("resume delivery status changed") + expected = _delivery_round_refs( + prior, seen=(*seen, str(prior_path)), + ) + snapshot = list(resume.get("round_record_refs", ())) + if not snapshot: + raise RuntimeError("resumed delivery lacks frozen prior-round refs") + if list(map(_ref_identity, snapshot)) != list(map(_ref_identity, expected)): + raise RuntimeError("resume prior-round snapshot changed") + prior_refs = snapshot + return [*prior_refs, *current_refs] + + +def _authenticated_round_records(delivery: dict) -> list[dict]: + refs = _delivery_round_refs(delivery) + records = [] + for ref in refs: + path = Path(ref["path"]).resolve() + if FUNNEL.sha256_file(path) != ref["sha256"]: + raise RuntimeError("round marker digest mismatch") + record = _read(path) + if record.get("status") != "SFM_B1_NEUTRAL_MULTIROUND_ROUND_COMPLETE": + raise RuntimeError("round marker status changed") + if "round" in ref and int(ref["round"]) != int(record["round"]): + raise RuntimeError("round marker index changed") + checkpoint = Path(record["checkpoint"]).resolve() + checkpoint_sha = record["checkpoint_sha256"] + if FUNNEL.sha256_file(checkpoint) != checkpoint_sha: + raise RuntimeError("round checkpoint digest mismatch") + has_post_d0 = "post_D0" in ref or "post_D0_sha256" in ref + if has_post_d0 and ( + str(Path(ref.get("post_D0", "")).resolve()) != str(checkpoint) + or ref.get("post_D0_sha256") != checkpoint_sha + ): + raise RuntimeError("round post-D0 snapshot changed") + has_post_positive = ( + "post_Dplus" in ref or "post_Dplus_sha256" in ref + ) + if has_post_positive: + post_positive = Path(ref.get("post_Dplus", "")).resolve() + if ( + not ref.get("post_Dplus_sha256") + or FUNNEL.sha256_file(post_positive) + != ref["post_Dplus_sha256"] + ): + raise RuntimeError("round post-Dplus digest mismatch") + records.append(record) + return records + + +def _validate_banks( + payloads: list[dict], *, screen_ep0: int, screen_M: int, + validation_ep0: int, validation_M: int, final_ep0: int, + expected_final_round: int, +) -> set[int]: + training_scenarios = set() + for payload in payloads: + screen = payload.get("disjoint_raw_evaluation", {}) + if ( + int(screen.get("ep0", -1)) != int(screen_ep0) + or int(screen.get("M_per_gamma", -1)) != int(screen_M) + ): + raise RuntimeError("delivery screening bank differs from declaration") + if int(payload.get("rounds", -1)) != int(expected_final_round): + raise RuntimeError("delivery does not reach the expected final round") + lineage = _authenticated_round_records(payload) + lineage_rounds = [int(record["round"]) for record in lineage] + if lineage_rounds != list(range(1, int(payload["rounds"]) + 1)): + raise RuntimeError("training round lineage is not contiguous") + arm_scenarios = { + int(scenario) for record in lineage + for scenario in record.get("scenarios", ()) + } + expected_scenarios = 2 * int(payload["rounds"]) + if len(arm_scenarios) != expected_scenarios: + raise RuntimeError("training scenario lineage is incomplete") + if training_scenarios and arm_scenarios != training_scenarios: + raise RuntimeError("arms used different training scenarios") + training_scenarios = arm_scenarios + banks = { + "screen": set(range(int(screen_ep0), int(screen_ep0) + int(screen_M))), + "validation": set(range( + int(validation_ep0), int(validation_ep0) + int(validation_M) + )), + "final": set(range(int(final_ep0), int(final_ep0) + 50)), + } + names = list(banks) + for index, name in enumerate(names): + if banks[name] & training_scenarios: + raise RuntimeError(f"{name} bank overlaps training scenarios") + for other in names[index + 1:]: + if banks[name] & banks[other]: + raise RuntimeError(f"{name} and {other} banks overlap") + return training_scenarios + + +def _mean(value: dict) -> float: + result = value.get("mean") + return float("nan") if result is None else float(result) + + +def _pooled(source: dict) -> dict: + return { + "SR": float(source["SR"]), + "CR": float(source["CR"]), + "timeout": float(source["timeout"]), + "Validity": _mean(source["Validity"]), + "clearance": _mean(source["successful_clearance"]), + "time_to_goal": _mean(source["successful_time_to_goal"]), + } + + +def _record_from_cell(payload: dict, *, method: str, temperature: float) -> dict: + record = payload["records"][0] + return { + "method": method, + "round": int(record["round"]), + "checkpoint": record["cell"]["checkpoint"], + "checkpoint_sha256": record["cell"]["checkpoint_sha256"], + "temperature": float(temperature), + "pooled": _pooled(record["cell"]["summary"]["pooled"]), + "per_gamma": { + gamma: _pooled(cell) + for gamma, cell in record["cell"]["summary"]["per_gamma"].items() + }, + "metrics_json": payload["metrics_json"] if "metrics_json" in payload else None, + } + + +def _screen_rows(root: Path, arm_names: list[str], payloads: list[dict]) -> tuple[list[dict], dict]: + rows, r0 = [], None + for arm_name, delivery in zip(arm_names, payloads): + lineage = { + int(record["round"]): record + for record in _authenticated_round_records(delivery) + } + for record in delivery["disjoint_raw_evaluation"]["records"]: + round_index = int(record["round"]) + if round_index == 0: + frozen_path = Path(delivery["checkpoint"]).resolve() + frozen_sha = delivery["checkpoint_sha256"] + else: + if round_index not in lineage: + raise RuntimeError("screening round is absent from lineage") + frozen_path = Path(lineage[round_index]["checkpoint"]).resolve() + frozen_sha = lineage[round_index]["checkpoint_sha256"] + checkpoint = Path( + record.get("checkpoint", frozen_path) + ).resolve() + expected_sha = record.get("checkpoint_sha256", frozen_sha) + if ( + checkpoint != frozen_path + or expected_sha != frozen_sha + or FUNNEL.sha256_file(checkpoint) != expected_sha + ): + raise RuntimeError("screening checkpoint digest mismatch") + pooled = _pooled(record["pooled"]) + row = { + "arm": arm_name, + "round": round_index, + "checkpoint": str(checkpoint), + "checkpoint_sha256": expected_sha, + "pooled": pooled, + } + if round_index == 0: + r0 = r0 or row + else: + rows.append(row) + if r0 is None: + raise RuntimeError("completed runs contain no r0 screening record") + return rows, r0 + + +def _screen_key(row: dict, r0: dict) -> tuple: + value, base = row["pooled"], r0["pooled"] + if any(not math.isfinite(float(value[key])) for key in METRICS): + return (1, float("inf"), float("inf"), float("inf"), float("inf"), + float("inf"), row["round"], row["arm"]) + wins = ( + int(value["CR"] < base["CR"]) + + int(value["Validity"] > base["Validity"]) + + int(value["clearance"] > base["clearance"]) + + int(value["time_to_goal"] < base["time_to_goal"]) + ) + return ( + 0, + -wins, + value["CR"], + -value["Validity"], + -value["clearance"], + value["time_to_goal"], + -value["SR"], + row["round"], + row["arm"], + ) + + +def _shortlist(rows: list[dict], r0: dict, arm_names: list[str]) -> list[dict]: + selected = [] + for arm in arm_names: + values = [row for row in rows if row["arm"] == arm] + if not values: + raise RuntimeError(f"no screening rows for {arm}") + selected.append(min(values, key=lambda row: _screen_key(row, r0))) + return selected + + +def _run_jobs(jobs: list[dict], *, gpus: list[int], workers: int, log_dir: Path) -> list[dict]: + queues = [[] for _ in gpus] + for index, job in enumerate(jobs): + queues[index % len(gpus)].append(job) + + def consume(gpu: int, cpu_start: int, values: list[dict]) -> list[dict]: + completed = [] + for job in values: + completed.append(FUNNEL.run_job( + job["name"], + job["command"], + gpu_index=gpu, + cpu_start=cpu_start, + cpu_count=workers, + log_dir=log_dir, + )) + print(f"COMPLETE {job['name']} gpu={gpu}", flush=True) + return completed + + with ThreadPoolExecutor(max_workers=len(gpus)) as executor: + futures = [ + executor.submit(consume, gpu, 16 + index * 80, queue) + for index, (gpu, queue) in enumerate(zip(gpus, queues)) + ] + return [item for future in futures for item in future.result()] + + +def _raw_command( + checkpoint: str, + *, + round_index: int, + temperature: float, + ep0: int, + noise_seed: int, + M: int, + workers: int, + output: Path, + cache: Path, +) -> list[str]: + return FUNNEL.evaluator_command( + [checkpoint], + [f"r{round_index}"], + scene_profile="double_density_velocity_ood", + ep0=ep0, + noise_seed=noise_seed, + m_per_gamma=M, + workers=workers, + cache_dir=cache, + output_dir=output, + temperature=temperature, + ) + + +def _compact_kazuki(rollout: dict, episode: int, gamma: float) -> dict: + status = ( + "success" if rollout["success"] + else "collision" if rollout["collision"] + else "timeout" + ) + steps = int(rollout["steps"]) + return { + "episode": int(episode), + "gamma": float(gamma), + "status": status, + "success": bool(rollout["success"]), + "collision": bool(rollout["collision"]), + "timeout": status == "timeout", + "steps": steps, + "time_to_goal": steps * SS.DT if rollout["success"] else None, + "min_clearance": float(rollout["min_clear"]), + "successful_clearance": ( + float(rollout["min_clear"]) if rollout["success"] else None + ), + "states": np.asarray(rollout["states"], np.float32), + "controls": np.asarray(rollout["controls"], np.float32), + "ped_xy": np.asarray(rollout["peds"], np.float32), + "ped_vel": np.asarray(rollout["ped_vels"], np.float32), + } + + +def run_kazuki(args) -> dict: + _source_gate(getattr(args, "expected_source_commit", None)) + policy, _ = GPS.load_sfm_policy(args.checkpoint, device=args.device) + policy.eval() + config = KZ.KazukiConfig(safe_coefs=(0.3,), goal_coef=0.5).validate() + environment = SS.scene_profile("double_density_velocity_ood") + rows = [] + for gamma in SP.GAMMAS: + for episode in range(int(args.ep0), int(args.ep0) + int(args.M)): + rollout = KZ.kazuki_sfm_deploy( + policy, + episode, + gamma, + cfg=config, + n_ped=environment["n_ped"], + T=SP.T, + device=args.device, + ped_speed_range=tuple(environment["ped_speed_range"]), + sample_seed=700_000, + collect_diagnostics=False, + ) + rows.append(_compact_kazuki(rollout, episode, gamma)) + del policy + import multiprocessing as mp + with ProcessPoolExecutor( + max_workers=int(args.workers), mp_context=mp.get_context("spawn") + ) as executor: + compact = EVAL._attach_validity(rows, executor) + summary = EVAL.summarize(compact, seed=int(args.ep0) + 700) + EVAL._assert_zero_verifier_errors(summary) + result = { + "status": "SFM_LOCKED_KAZUKI_EXACT_VALIDITY_COMPLETE", + "method": "locked Kazuki generate-guide-refine", + "safe_coef": 0.3, + "goal_coef": 0.5, + "checkpoint": str(Path(args.checkpoint).resolve()), + "checkpoint_sha256": FUNNEL.sha256_file(Path(args.checkpoint)), + "scene_profile": "double_density_velocity_ood", + "ep0": int(args.ep0), + "M_per_gamma": int(args.M), + "summary": summary, + "rows": compact, + } + _write(Path(args.output), result) + return result + + +def _kazuki_record(payload: dict) -> dict: + return { + "method": "kazuki_locked", + "round": 0, + "checkpoint": payload["checkpoint"], + "checkpoint_sha256": payload["checkpoint_sha256"], + "temperature": None, + "pooled": _pooled(payload["summary"]["pooled"]), + "per_gamma": { + gamma: _pooled(cell) + for gamma, cell in payload["summary"]["per_gamma"].items() + }, + } + + +def _envelope(pretrained: list[dict], kazuki: dict) -> dict: + references = [*pretrained, kazuki] + return { + metric: ( + min(row["pooled"][metric] for row in references) + if metric in LOWER_IS_BETTER + else max(row["pooled"][metric] for row in references) + ) + for metric in METRICS + } + + +def _shortfalls(row: dict, target: dict) -> dict: + value = row["pooled"] + result = {} + for metric in METRICS: + if not math.isfinite(float(value[metric])): + result[metric] = float("inf") + continue + scale = max(abs(float(target[metric])), .02) + if metric in LOWER_IS_BETTER: + result[metric] = max(0.0, value[metric] - target[metric]) / scale + else: + result[metric] = max(0.0, target[metric] - value[metric]) / scale + return result + + +def _liveness_contract(pretrained: list[dict], kazuki: dict) -> dict: + references = [*pretrained, kazuki] + return { + "minimum_SR": max( + 0.0, min(row["pooled"]["SR"] for row in references) - .05 + ), + "maximum_timeout": min( + 1.0, max(row["pooled"]["timeout"] for row in references) + .05 + ), + "every_gamma_has_success": True, + } + + +def _liveness_eligible(row: dict, contract: dict) -> bool: + if ( + row["pooled"]["SR"] < contract["minimum_SR"] + or row["pooled"]["timeout"] > contract["maximum_timeout"] + ): + return False + return all( + math.isfinite(float(cell["clearance"])) + and math.isfinite(float(cell["time_to_goal"])) + for cell in row["per_gamma"].values() + ) + + +def _trend(record: dict) -> dict: + rows = [record["per_gamma"][str(gamma)] for gamma in SP.GAMMAS] + tests = { + "CR_low_gamma_not_higher": [ + rows[i]["CR"] <= rows[i + 1]["CR"] + .10 + for i in range(len(rows) - 1) + ], + "Validity_rises_with_gamma": [ + rows[i]["Validity"] <= rows[i + 1]["Validity"] + .10 + for i in range(len(rows) - 1) + ], + "clearance_falls_with_gamma": [ + rows[i]["clearance"] >= rows[i + 1]["clearance"] - .02 + for i in range(len(rows) - 1) + ], + "time_falls_with_gamma": [ + rows[i]["time_to_goal"] >= rows[i + 1]["time_to_goal"] - 1.0 + for i in range(len(rows) - 1) + ], + } + fractions = { + name: sum(values) / len(values) for name, values in tests.items() + } + return { + "adjacent_pair_fractions": fractions, + "mean_fraction": float(np.mean(list(fractions.values()))), + "tolerances": {"rate": .10, "clearance_m": .02, "time_s": 1.0}, + } + + +def _trend_eligible(record: dict) -> bool: + """Require every requested gamma-ordering family, not just its mean.""" + fractions = _trend(record)["adjacent_pair_fractions"] + return all( + value >= MIN_TREND_FAMILY_FRACTION for value in fractions.values() + ) + + +def _selection_key(row: dict, target: dict, liveness: dict) -> tuple: + shortfall = _shortfalls(row, target) + trend = _trend(row) + return ( + 0 if _liveness_eligible(row, liveness) else 1, + max(shortfall.values()), + sum(shortfall.values()), + -trend["mean_fraction"], + -row["pooled"]["SR"], + row["pooled"]["timeout"], + row["round"], + row["method"], + row["temperature"], + ) + + +def _render(records: list[dict], output: Path) -> None: + specs = ( + ("CR", "Collision rate"), + ("Validity", "Validity"), + ("clearance", "Min. clearance [m]"), + ("time_to_goal", "Time-to-goal [s]"), + ) + colors = { + "pretrained": "#7f7f7f", + "expanded": "#0072B2", + "kazuki_locked": "#CC79A7", + } + figure, axes = plt.subplots(2, 2, figsize=(14.6, 10.8)) + gammas = np.asarray(SP.GAMMAS, float) + for axis, (metric, title) in zip(axes.flat, specs): + for record in records: + values = [ + record["per_gamma"][str(gamma)][metric] for gamma in SP.GAMMAS + ] + label = record["method"] + if record["temperature"] is not None: + label += f" (temp={record['temperature']:g})" + axis.plot( + gammas, values, marker="o", lw=2.4, + color=colors[record["method"]], label=label, + ) + axis.set_title(title) + axis.set_xlabel(r"$\gamma$") + axis.grid(alpha=.25) + if metric in {"CR", "Validity"}: + axis.set_ylim(-.03, 1.03) + handles, labels = axes[0, 0].get_legend_handles_labels() + figure.legend(handles, labels, loc="upper center", ncol=3, frameon=False) + figure.tight_layout(rect=(0, 0, 1, .93)) + output.parent.mkdir(parents=True, exist_ok=True) + figure.savefig(output, dpi=300, bbox_inches="tight") + figure.savefig(output.with_suffix(".pdf"), bbox_inches="tight") + plt.close(figure) + + +def run(args) -> dict: + started = datetime.now(timezone.utc) + source_commit = _source_gate( + getattr(args, "expected_source_commit", None) + ) + training_root = Path(args.training_root).resolve() + output = Path(args.output_dir).resolve() + if output.exists(): + raise FileExistsError(output) + output.mkdir(parents=True) + arm_names = args.arm_names.split(",") + temperatures = [float(item) for item in args.temperatures.split(",")] + if 1.0 not in temperatures or any(item <= 0 or item > 1 for item in temperatures): + raise ValueError("temperatures must be in (0,1] and include 1") + deliveries = _wait_for_deliveries(training_root, arm_names, args.poll_seconds) + training_scenarios = _validate_banks( + deliveries, + screen_ep0=args.screen_ep0, + screen_M=args.screen_M, + validation_ep0=args.validation_ep0, + validation_M=args.validation_M, + final_ep0=args.final_ep0, + expected_final_round=args.expected_final_round, + ) + screen_rows, screen_r0 = _screen_rows( + training_root, arm_names, deliveries + ) + shortlist = _shortlist(screen_rows, screen_r0, arm_names) + pretrained = str(Path(deliveries[0]["checkpoint"]).resolve()) + + validation_root = output / "validation_m10" + cache = validation_root / "cache" + jobs = [] + job_meta = {} + for temperature in temperatures: + name = f"pretrained_temp{temperature:g}".replace(".", "p") + out = validation_root / name + jobs.append({ + "name": name, + "command": _raw_command( + pretrained, round_index=0, temperature=temperature, + ep0=args.validation_ep0, noise_seed=args.validation_noise_seed, + M=args.validation_M, workers=args.workers, + output=out, cache=cache, + ), + }) + job_meta[name] = ("pretrained", 0, temperature, out) + for candidate in shortlist: + for temperature in temperatures: + name = ( + f"{candidate['arm']}_r{candidate['round']}_temp{temperature:g}" + ).replace(".", "p") + out = validation_root / name + jobs.append({ + "name": name, + "command": _raw_command( + candidate["checkpoint"], round_index=candidate["round"], + temperature=temperature, ep0=args.validation_ep0, + noise_seed=args.validation_noise_seed, + M=args.validation_M, workers=args.workers, + output=out, cache=cache, + ), + }) + job_meta[name] = ( + candidate["arm"], candidate["round"], temperature, out + ) + _run_jobs( + jobs, gpus=args.gpus, workers=args.workers, + log_dir=output / "logs", + ) + + validation_records = [] + for name, (method, round_index, temperature, out) in job_meta.items(): + payload = _read(out / f"raw_m{args.validation_M}_offline_metrics.json") + payload["metrics_json"] = str( + out / f"raw_m{args.validation_M}_offline_metrics.json" + ) + validation_records.append(_record_from_cell( + payload, method=method, temperature=temperature + )) + + kazuki_validation_path = validation_root / "kazuki_locked.json" + kazuki_command = [ + sys.executable, str(Path(__file__).resolve()), "kazuki", + "--checkpoint", pretrained, + "--ep0", str(args.validation_ep0), + "--M", str(args.validation_M), + "--workers", str(args.workers), + "--device", "cuda:0", + "--output", str(kazuki_validation_path), + ] + FUNNEL.run_job( + "kazuki_validation", + kazuki_command, + gpu_index=args.gpus[0], + cpu_start=16, + cpu_count=args.workers, + log_dir=output / "logs", + ) + kazuki_validation = _kazuki_record(_read(kazuki_validation_path)) + pretrained_records = [ + row for row in validation_records if row["method"] == "pretrained" + ] + expanded_records = [ + row for row in validation_records if row["method"] != "pretrained" + ] + target = _envelope(pretrained_records, kazuki_validation) + liveness = _liveness_contract(pretrained_records, kazuki_validation) + selected_expanded = min( + expanded_records, + key=lambda row: _selection_key(row, target, liveness), + ) + selected_pretrained = min( + pretrained_records, + key=lambda row: _selection_key(row, target, liveness), + ) + selection = { + "status": "LOCKED_BEFORE_DISJOINT_M50", + "rule": ( + "minimize maximum then total normalized shortfall from the " + "metric-wise envelope of all pretrained temperatures and locked " + "Kazuki; break ties by gamma-trend score, SR, timeout, round, name" + ), + "global_temperature_only": True, + "per_gamma_temperature_forbidden": True, + "target_envelope": target, + "liveness_contract": liveness, + "selected_expanded": selected_expanded, + "selected_pretrained": selected_pretrained, + "expanded_shortfalls": _shortfalls(selected_expanded, target), + "expanded_point_estimate_four_metric_gate": ( + _liveness_eligible(selected_expanded, liveness) + and all( + value == 0 + for value in _shortfalls(selected_expanded, target).values() + ) + ), + "expanded_gamma_trend": _trend(selected_expanded), + } + _write(output / "SELECTION_LOCKED.json", selection) + + final_root = output / "disjoint_m50" + final_cache = final_root / "cache" + final_jobs = [] + final_meta = {} + for method, record in ( + ("pretrained", selected_pretrained), + ("expanded", selected_expanded), + ): + out = final_root / method + final_jobs.append({ + "name": f"final_{method}", + "command": _raw_command( + record["checkpoint"], round_index=record["round"], + temperature=record["temperature"], ep0=args.final_ep0, + noise_seed=args.final_noise_seed, M=50, + workers=args.workers, output=out, cache=final_cache, + ), + }) + final_meta[method] = (record, out) + final_kazuki = final_root / "kazuki_locked.json" + final_jobs.append({ + "name": "final_kazuki", + "command": [ + sys.executable, str(Path(__file__).resolve()), "kazuki", + "--checkpoint", pretrained, + "--ep0", str(args.final_ep0), + "--M", "50", + "--workers", str(args.workers), + "--device", "cuda:0", + "--output", str(final_kazuki), + ], + }) + _run_jobs( + final_jobs, gpus=args.gpus, workers=args.workers, + log_dir=output / "logs", + ) + + final_records = [] + for method, (record, out) in final_meta.items(): + payload = _read(out / "raw_m50_offline_metrics.json") + final_records.append(_record_from_cell( + payload, method=method, temperature=record["temperature"] + )) + final_records.append(_kazuki_record(_read(final_kazuki))) + expanded_final = next(row for row in final_records if row["method"] == "expanded") + references = [ + row for row in final_records if row["method"] != "expanded" + ] + final_target = _envelope( + [row for row in references if row["method"] == "pretrained"], + next(row for row in references if row["method"] == "kazuki_locked"), + ) + final_shortfall = _shortfalls(expanded_final, final_target) + _render(final_records, final_root / "four_metric_per_gamma.png") + result = { + "status": STATUS, + "source_commit": source_commit, + "training_root": str(training_root), + "training_scenario_ids": sorted(training_scenarios), + "training_deliveries": [ + str(training_root / name / "DELIVERY_COMPLETE.json") + for name in arm_names + ], + "banks": { + "training_screen": { + "M_per_gamma": args.screen_M, "ep0": args.screen_ep0 + }, + "temperature_validation": { + "M_per_gamma": args.validation_M, + "ep0": args.validation_ep0, + "noise_seed": args.validation_noise_seed, + }, + "disjoint_confirmation": { + "M_per_gamma": 50, + "ep0": args.final_ep0, + "noise_seed": args.final_noise_seed, + }, + }, + "temperatures": temperatures, + "shortlist": shortlist, + "selection": selection, + "final_records": final_records, + "final_target_envelope": final_target, + "final_expanded_shortfalls": final_shortfall, + "final_point_estimate_four_metric_gate": all( + value == 0 for value in final_shortfall.values() + ), + "scientific_win_status": ( + "PENDING_PAIRED_SCENARIO_CLUSTER_CI; point estimates alone are " + "not a win claim" + ), + "final_gamma_trends": { + row["method"]: _trend(row) for row in final_records + }, + "completed_at": datetime.now(timezone.utc).isoformat(), + "wall_seconds": ( + datetime.now(timezone.utc) - started + ).total_seconds(), + } + _source_gate(source_commit) + _write(output / "DELIVERY_COMPLETE.json", result) + print(json.dumps({ + "status": STATUS, + "point_estimate_four_metric_gate": ( + result["final_point_estimate_four_metric_gate"] + ), + "expanded": expanded_final["pooled"], + "delivery": str(output / "DELIVERY_COMPLETE.json"), + }, indent=2)) + return result + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(description=__doc__) + sub = parser.add_subparsers(dest="command", required=True) + full = sub.add_parser("run") + full.add_argument("--training-root", required=True) + full.add_argument("--output-dir", required=True) + full.add_argument( + "--arm-names", + default="lr1em5_s01,lr1em5_s04,lr3em5_s01,lr3em5_s04", + ) + full.add_argument("--temperatures", default="0.55,0.7,0.85,1.0") + full.add_argument("--gpus", type=int, nargs="+", default=[1, 3]) + full.add_argument("--workers", type=int, default=32) + full.add_argument("--poll-seconds", type=int, default=60) + full.add_argument("--screen-ep0", type=int, default=270_000) + full.add_argument("--screen-M", type=int, default=20) + full.add_argument("--validation-ep0", type=int, default=460_000) + full.add_argument("--validation-M", type=int, default=10) + full.add_argument("--validation-noise-seed", type=int, default=2_026_073_4) + full.add_argument("--final-ep0", type=int, default=470_000) + full.add_argument("--final-noise-seed", type=int, default=2_026_073_5) + full.add_argument("--expected-final-round", type=int, default=50) + full.add_argument("--expected-source-commit") + + kazuki = sub.add_parser("kazuki") + kazuki.add_argument("--checkpoint", required=True) + kazuki.add_argument("--ep0", type=int, required=True) + kazuki.add_argument("--M", type=int, required=True) + kazuki.add_argument("--workers", type=int, default=32) + kazuki.add_argument("--device", default="cuda:0") + kazuki.add_argument("--output", required=True) + kazuki.add_argument("--expected-source-commit") + return parser + + +def main(argv=None) -> int: + args = build_parser().parse_args(argv) + if args.command == "kazuki": + run_kazuki(args) + else: + run(args) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/overnight_run_07_12_sfm/sfm_b1_branch_compare_viz.py b/overnight_run_07_12_sfm/sfm_b1_branch_compare_viz.py new file mode 100644 index 0000000..96a514f --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_branch_compare_viz.py @@ -0,0 +1,163 @@ +"""Compare final planned-window branch forests for two SFM checkpoints.""" +from __future__ import annotations + +import argparse +import json +import os + +import matplotlib +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import torch + +import sfm_b1_d_branch_viz as DB +import sfm_b1_full_episode_viz as FV + + +STATUS = "SFM_B1_BRANCH_COMPARISON_COMPLETE" + + +def summarize(bundle): + traces = list(bundle["traces"]) + outcomes = list(bundle["outcomes"]) + positive = sum( + row["executed_label"] == "verifier_positive" for row in traces + ) + negative = sum( + row["executed_label"] == "verifier_negative" for row in traces + ) + resolved = positive + negative + return { + "contexts": len(traces), + "executed_positive": positive, + "executed_negative": negative, + "executed_positive_fraction": ( + float(positive / resolved) if resolved else None + ), + "success": sum(bool(row["success"]) for row in outcomes), + "collision": sum(bool(row["collision"]) for row in outcomes), + "timeout": sum(bool(row["timeout"]) for row in outcomes), + "episodes": len(outcomes), + } + + +def _load(path): + bundle = torch.load(path, map_location="cpu", weights_only=False) + if bundle.get("status") != "SFM_B1_FULL_EPISODE_LABEL_AUDIT_COMPLETE": + raise ValueError(f"not a completed branch audit: {path}") + return bundle + + +def render( + pretrained_trace, expanded_trace, output_png, output_json, + *, expanded_label="partial expanded", +): + pretrained = _load(pretrained_trace) + expanded = _load(expanded_trace) + for key in ("scenarios", "gammas", "environment", "sample_seed", "audit_seed"): + if pretrained[key] != expanded[key]: + raise ValueError(f"comparison contract differs at {key}") + scenarios = tuple(map(int, pretrained["scenarios"])) + gammas = tuple(map(float, pretrained["gammas"])) + if len(scenarios) != 3 or len(gammas) != 7: + raise ValueError("comparison requires three scenarios and seven gammas") + + models = (("pretrained", pretrained), (expanded_label, expanded)) + figure, axes = plt.subplots(6, 7, figsize=(23.4, 18.0)) + figure.subplots_adjust( + left=.055, right=.82, bottom=.025, top=.96, wspace=.025, hspace=.04, + ) + for column, gamma in enumerate(gammas): + figure.text( + .055 + (.765 / 7) * (column + .5), .975, + f"$\\gamma={gamma:g}$", ha="center", va="center", fontsize=10, + ) + + reports = {} + for model_index, (label, bundle) in enumerate(models): + index = FV._index(bundle["traces"]) + reports[label] = summarize(bundle) + for scenario_index, scenario in enumerate(scenarios): + row = model_index * len(scenarios) + scenario_index + figure.text( + .018, .96 - (.935 / 6) * (row + .5), + f"{label}\nepisode {scenario}", + ha="center", va="center", rotation=90, fontsize=8, + ) + for column, gamma in enumerate(gammas): + rows = index[(scenario, round(gamma, 8))] + DB.draw_cell(axes[row, column], rows, max(rows)) + + figure.legend( + handles=DB._legend(), loc="center left", bbox_to_anchor=(.835, .66), + frameon=False, fontsize=8, + ) + text = [] + for label, report in reports.items(): + text.extend([ + label, + f"verified D: {report['executed_positive']}/" + f"{report['contexts']} " + f"({report['executed_positive_fraction']:.1%})", + f"outcomes S/C/T: {report['success']}/" + f"{report['collision']}/{report['timeout']}", + "", + ]) + text.extend([ + "Fixed scenarios, gamma values, and", + "proposal x0 streams are shared.", + "Contexts diverge after the first", + "closed-loop action.", + ]) + figure.text( + .835, .42, "\n".join(text), ha="left", va="top", fontsize=8, + ) + + os.makedirs(os.path.dirname(os.path.abspath(output_png)), exist_ok=True) + figure.savefig(output_png, dpi=165, bbox_inches="tight") + plt.close(figure) + report = { + "status": STATUS, + "pretrained_trace": os.path.abspath(pretrained_trace), + "expanded_trace": os.path.abspath(expanded_trace), + "expanded_label": expanded_label, + "comparison_contract": { + "scenarios": list(scenarios), + "gammas": list(gammas), + "sample_seed": pretrained["sample_seed"], + "audit_seed": pretrained["audit_seed"], + "closed_loop_caveat": ( + "proposal x0 streams match by cell and step, but checkpoint-" + "dependent first actions make later contexts different" + ), + }, + "models": reports, + "png": os.path.abspath(output_png), + } + os.makedirs(os.path.dirname(os.path.abspath(output_json)), exist_ok=True) + temporary = os.path.abspath(output_json) + ".tmp" + with open(temporary, "w") as stream: + json.dump(report, stream, indent=2, allow_nan=False) + os.replace(temporary, os.path.abspath(output_json)) + return report + + +def main(argv=None): + parser = argparse.ArgumentParser() + parser.add_argument("--pretrained-trace", required=True) + parser.add_argument("--expanded-trace", required=True) + parser.add_argument("--expanded-label", default="partial expanded") + parser.add_argument("--output-png", required=True) + parser.add_argument("--output-json", required=True) + args = parser.parse_args(argv) + render( + args.pretrained_trace, + args.expanded_trace, + args.output_png, + args.output_json, + expanded_label=args.expanded_label, + ) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/sfm_b1_cost.py b/overnight_run_07_12_sfm/sfm_b1_cost.py index 8821fbf..eed5db4 100644 --- a/overnight_run_07_12_sfm/sfm_b1_cost.py +++ b/overnight_run_07_12_sfm/sfm_b1_cost.py @@ -13,6 +13,9 @@ import sfm_scene as SS +PROGRESS_GATE_DISPLACEMENT_M = 0.2 + + def frozen_expert_config(): """The exact mode-1 adapter config whose cost terms define arms B--D.""" values = GS.mode1_config(range_m=SS.R_SENSE, u_max=SS.U_MAX, noise_var_mult=3.0) @@ -95,6 +98,42 @@ def nominal_hp_margin(state, first_action, ped_xy, gamma): return new - (1.0 - float(gamma)) * old, old, new +def select_progress_gated_margin(rows, *, state): + """Prefer progressive, non-stalling H10 plans, then maximize Hp margin. + + This is a selector only. The caller remains responsible for the exact + verifier and nominal-Hp admission gates. If no admitted row clears the + progress gate, the historical max-margin rule is used unchanged. + """ + if not rows: + return None + initial = np.asarray(state, np.float32).reshape(4)[:2] + goal_distance = float(np.linalg.norm(initial - SS.GOAL)) + progressive = [] + for row in rows: + states = rollout_states( + state, + np.asarray(row["controls"], np.float32)[None], + )[0].detach().cpu().numpy() + displacement = float(np.linalg.norm(states[-1, :2] - initial)) + progress = goal_distance - float( + np.linalg.norm(states[-1, :2] - SS.GOAL) + ) + row["H10_displacement"] = displacement + row["H10_goal_progress"] = progress + row["progress_gate_eligible"] = bool( + displacement >= PROGRESS_GATE_DISPLACEMENT_M + and progress > 0.0 + ) + if row["progress_gate_eligible"]: + progressive.append(row) + support = progressive or rows + return max( + support, + key=lambda row: (row["hp_margin"], -int(row["candidate_id"])), + ) + + def select_admissible(query_rows, *, selector, state, ped_xy, ped_vel, gamma): """Gate by y and one-step nominal Hp before either frozen selector.""" admissible = [] @@ -111,13 +150,40 @@ def select_admissible(query_rows, *, selector, state, ped_xy, ped_vel, gamma): return None if selector == "margin": return max(admissible, key=lambda row: (row["hp_margin"], -int(row["candidate_id"]))) - if selector != "safemppi_cost": + if selector == "progress_gated_margin": + return select_progress_gated_margin(admissible, state=state) + if selector not in ("safemppi_cost", "balanced_rank"): raise ValueError(f"unknown selector: {selector}") controls = torch.as_tensor(np.stack([row["controls"] for row in admissible]), dtype=torch.float32) costs = safemppi_proposal_cost(state, controls, SS.GOAL, ped_xy, ped_vel).cpu().numpy() for row, cost in zip(admissible, costs): row["expert_cost"] = float(cost) - return min(admissible, key=lambda row: (row["expert_cost"], int(row["candidate_id"]))) + if selector == "safemppi_cost": + return min(admissible, key=lambda row: (row["expert_cost"], int(row["candidate_id"]))) + safety_order = sorted( + admissible, + key=lambda row: (-row["hp_margin"], int(row["candidate_id"])), + ) + performance_order = sorted( + admissible, + key=lambda row: (row["expert_cost"], int(row["candidate_id"])), + ) + for rank, row in enumerate(safety_order, start=1): + row["safety_rank"] = rank + for rank, row in enumerate(performance_order, start=1): + row["performance_rank"] = rank + for row in admissible: + row["rank_sum"] = row["safety_rank"] + row["performance_rank"] + return min( + admissible, + key=lambda row: ( + row["rank_sum"], + row["safety_rank"], + -row["hp_margin"], + row["expert_cost"], + int(row["candidate_id"]), + ), + ) def scorer_manifest(): diff --git a/overnight_run_07_12_sfm/sfm_b1_d_branch_viz.py b/overnight_run_07_12_sfm/sfm_b1_d_branch_viz.py new file mode 100644 index 0000000..46edec9 --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_d_branch_viz.py @@ -0,0 +1,291 @@ +"""Render the one planned H=10 D sample attached to every executed context. + +Each thin blue/red branch is the exact H=10 window whose first action advanced +the offline collector at that context. The thick black path joins those first +actions. Current K/B queries and the current executed-window verifier geometry +remain visible so the branch origin can be audited. +""" +from __future__ import annotations + +import argparse +import json +import os + +import matplotlib +matplotlib.use("Agg") +import matplotlib.animation as animation +import matplotlib.pyplot as plt +from matplotlib.lines import Line2D +import numpy as np +import torch + +import _paths # noqa: F401 +import sfm_b1_density_viz as DV +import sfm_b1_full_episode_viz as FV +import sfm_b1_viz as BV +import sfm_metrics2 as SM +import sfm_scene as SS + + +def _branch(trace): + result = trace["executed_result"] + path = np.asarray(result.get("segment", ()), float) + if path.shape != (11, 2): + path = SM.rollout_positions( + np.asarray(trace["state"], float), + np.asarray(trace["executed_controls"], float), + ) + return path + + +def _draw_D_branches(axis, rows, step, *, line_scale=1.0): + available = sorted(value for value in rows if value <= int(step)) + for context_step in available: + trace = rows[context_step] + path = _branch(trace) + color = FV._executed_color(trace) + is_current = context_step == available[-1] + axis.plot( + path[:, 0], path[:, 1], + color=color, + lw=(1.25 if is_current else .55) * float(line_scale), + marker=".", + ms=(1.3 if is_current else .75) * np.sqrt(float(line_scale)), + alpha=.9 if is_current else .32, + zorder=7 if is_current else 3, + ) + axis.plot( + path[0, 0], path[0, 1], marker=".", color=color, + ms=(2.7 if is_current else 1.5) * np.sqrt(float(line_scale)), + zorder=8, + ) + + +def _draw_executed_trajectory( + axis, rows, step, *, linewidth=2.8, marker_size=2.1, +): + available = sorted(value for value in rows if value <= int(step)) + if not available: + return + states = [np.asarray(rows[value]["state"], float)[:2] for value in available] + states.append(np.asarray(rows[available[-1]]["next_state"], float)[:2]) + states = np.asarray(states) + axis.plot( + states[:, 0], states[:, 1], color="#111111", lw=float(linewidth), + marker=".", ms=float(marker_size), alpha=.97, zorder=11, + ) + axis.annotate( + "", xy=states[-1], xytext=states[-2], + arrowprops=dict( + arrowstyle="->", color="#111111", + lw=max(.8, .78 * float(linewidth)), + ), + zorder=12, + ) + + +def _robot_frame(path, trace): + position = np.asarray(trace["state"], float)[:2] + direction = np.asarray(trace["state"], float)[2:4] + if np.linalg.norm(direction) < 1.0e-6: + direction = np.asarray(SS.GOAL, float)[:2] - position + direction = direction / max(np.linalg.norm(direction), 1.0e-12) + normal = np.array([-direction[1], direction[0]]) + delta = np.asarray(path, float) - position + return np.stack([delta @ direction, delta @ normal], axis=1) + + +def _draw_candidate_inset(axis, trace): + for child in list(axis.child_axes): + if getattr(child, "_sfm_candidate_inset", False): + child.remove() + inset = axis.inset_axes((.035, .675, .30, .29), zorder=30) + inset._sfm_candidate_inset = True + inset.set_facecolor((1., 1., 1., .68)) + selected_id = trace.get("executed_id") + local_paths = [] + for query_index, row in enumerate(trace["query_rows"], start=1): + path = np.asarray( + BV._trace_candidate(trace, int(row["candidate_id"]))["segment"], + float, + ) + local = _robot_frame(path, trace) + local_paths.append(local) + status, _ = BV._candidate_status(trace, int(row["candidate_id"])) + color = ( + BV.GREEN if status == "positive" + else BV.RED if status == "negative" + else BV.GRAY + ) + is_selected = ( + selected_id is not None + and int(row["candidate_id"]) == int(selected_id) + ) + if is_selected: + inset.plot( + local[:, 0], local[:, 1], color="#111111", + lw=3.3, alpha=.82, zorder=3, + ) + inset.plot( + local[:, 0], local[:, 1], color=color, + lw=2.25 if is_selected else 1.05, + alpha=.98 if is_selected else .72, + zorder=4 if is_selected else 2, + ) + inset.text( + local[-1, 0], local[-1, 1], str(query_index), + fontsize=4.8, color=color, ha="center", va="center", zorder=5, + ) + inset.plot(0., 0., marker=">", color="#111111", ms=3.2, zorder=6) + if local_paths: + joined = np.concatenate(local_paths) + span = max(.18, 1.08 * float(np.max(np.abs(joined)))) + inset.set_xlim(-.08 * span, span) + inset.set_ylim(-span, span) + inset.set_aspect("equal") + inset.set_xticks([]) + inset.set_yticks([]) + inset.set_title( + "robot-frame B=4" if selected_id is not None else "robot-frame B=4 · NVP", + fontsize=5.1, pad=1.2, + ) + for spine in inset.spines.values(): + spine.set_alpha(.34) + spine.set_linewidth(.55) + + +def draw_cell( + axis, rows, step, *, branch_line_scale=1.0, + trajectory_linewidth=2.8, trajectory_marker_size=2.1, + candidate_inset=False, +): + available = [value for value in rows if value <= int(step)] + current_step = max(available) if available else min(rows) + trace = rows[current_step] + BV._draw_common(axis, trace, nominal_levels=False) + _draw_D_branches( + axis, rows, current_step, line_scale=branch_line_scale, + ) + FV._draw_candidates(axis, trace) + FV._draw_executed(axis, trace) + _draw_executed_trajectory( + axis, rows, current_step, linewidth=trajectory_linewidth, + marker_size=trajectory_marker_size, + ) + if candidate_inset: + _draw_candidate_inset(axis, trace) + DV._set_clean_axis(axis) + return trace + + +def _legend(): + return [ + Line2D([], [], color=BV.BLUE, lw=1.2, label=r"$D^+$ planned H10 branch"), + Line2D([], [], color=BV.RED, lw=1.2, label=r"$D^-$ planned H10 branch"), + Line2D([], [], color="#111111", lw=2.8, label="executed first-action trajectory"), + Line2D([], [], color=BV.GRAY, lw=.7, label="current K=16 generated"), + Line2D([], [], color=BV.ORANGE, lw=1.1, label="current B=4 RBF queried"), + Line2D([], [], color=BV.GREEN, lw=1.2, label="current B full-H positive"), + Line2D([], [], color=BV.RED, lw=1.2, marker="x", label="current B full-H rejected"), + Line2D([], [], color=BV.GREEN, lw=.7, label="executed verifier levels h=1..10"), + ] + + +def render(trace_path, output_mp4, output_png, output_json, *, fps=5, frame_stride=2): + bundle = torch.load(trace_path, map_location="cpu", weights_only=False) + if bundle.get("status") != "SFM_B1_FULL_EPISODE_LABEL_AUDIT_COMPLETE": + raise ValueError("input is not a completed full-episode audit") + scenarios = tuple(map(int, bundle["scenarios"])) + gammas = tuple(map(float, bundle["gammas"])) + if len(scenarios) != 3 or gammas != tuple(map(float, SS.GAMMAS)): + raise ValueError("renderer requires three scenarios and all seven gammas") + index = FV._index(bundle["traces"]) + maximum = max(max(rows) for rows in index.values()) + frames = list(range(0, maximum + 1, int(frame_stride))) + if frames[-1] != maximum: + frames.append(maximum) + + figure, axes = plt.subplots(3, 7, figsize=(23.2, 10.1)) + figure.subplots_adjust( + left=.035, right=.815, bottom=.025, top=.94, wspace=.025, hspace=.04, + ) + for column, gamma in enumerate(gammas): + figure.text( + .035 + (.78 / 7) * (column + .5), .965, f"$\\gamma={gamma:g}$", + ha="center", va="center", fontsize=10, + ) + for row, scenario in enumerate(scenarios): + figure.text( + .012, .94 - (.915 / 3) * (row + .5), f"episode\n{scenario}", + ha="center", va="center", rotation=90, fontsize=9, + ) + figure.legend( + handles=_legend(), loc="center left", bbox_to_anchor=(.825, .59), + frameon=False, fontsize=8, + ) + figure.text( + .825, .26, + "One D sample per context.\n" + "Each branch is the complete planned H10 window;\n" + "only its first action advances the black trajectory.\n" + "Finite-B NVP does not stop this offline collector.", + ha="left", va="top", fontsize=8, + ) + + def update(step): + for row, scenario in enumerate(scenarios): + for column, gamma in enumerate(gammas): + axis = axes[row, column] + axis.clear() + draw_cell(axis, index[(scenario, round(gamma, 8))], int(step)) + return [] + + for path in (output_mp4, output_png, output_json): + os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True) + movie = animation.FuncAnimation( + figure, update, frames=frames, interval=1000 / int(fps), blit=False, + ) + movie.save( + output_mp4, writer=animation.FFMpegWriter(fps=int(fps), bitrate=4600), + dpi=105, + ) + update(maximum) + figure.savefig(output_png, dpi=165, bbox_inches="tight") + plt.close(figure) + + report = dict( + status="SFM_B1_D_BRANCH_VIZ_COMPLETE", + trace_path=os.path.abspath(trace_path), + scenarios=list(scenarios), + gammas=list(gammas), + D_semantics=( + "one exact full-H10 planned window per executed context; blue y=1, " + "red y=0; thick black joins executed first actions" + ), + frames=frames, + mp4=os.path.abspath(output_mp4), + png=os.path.abspath(output_png), + ) + with open(output_json + ".tmp", "w") as stream: + json.dump(report, stream, indent=2) + os.replace(output_json + ".tmp", output_json) + return report + + +def main(argv=None): + parser = argparse.ArgumentParser() + parser.add_argument("--trace", required=True) + parser.add_argument("--output-mp4", required=True) + parser.add_argument("--output-png", required=True) + parser.add_argument("--output-json", required=True) + parser.add_argument("--fps", type=int, default=5) + parser.add_argument("--frame-stride", type=int, default=2) + args = parser.parse_args(argv) + render( + args.trace, args.output_mp4, args.output_png, args.output_json, + fps=args.fps, frame_stride=args.frame_stride, + ) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/sfm_b1_density_viz.py b/overnight_run_07_12_sfm/sfm_b1_density_viz.py index 5c3a487..f2a2616 100644 --- a/overnight_run_07_12_sfm/sfm_b1_density_viz.py +++ b/overnight_run_07_12_sfm/sfm_b1_density_viz.py @@ -37,6 +37,7 @@ "selected": "Arm-A r10 learned raw", "kazuki": "Kazuki generate-guide-refine", } +CYAN = "#00A6D6" MAGENTA = "#CC79A7" METHOD_ALIASES = { "expert": ("expert", "safemppi_expert"), @@ -574,22 +575,56 @@ def draw_method_panel(axis, method, run, gamma, step, *, verifier_result=None, axis.plot(plan[:, 0], plan[:, 1], color="#7F3C8D", lw=.82, marker="o", ms=2.2) guidance = trace.get("accumulated_guidance") if guidance: - vector = float(guidance_scale) * np.asarray(guidance["net_guidance_action"], float) - norm = float(np.linalg.norm(vector)) - if norm > float(guidance_cap): - vector *= float(guidance_cap) / norm start = np.asarray(trace["state"], float)[:2] - arrow = FancyArrowPatch( - tuple(start), tuple(start + vector), arrowstyle="-|>", mutation_scale=15, - lw=2.6, color=MAGENTA, shrinkA=0, shrinkB=0, zorder=11, - ) - axis.add_patch(arrow) - metadata.update( - guidance_present=True, - net_guidance_action=np.asarray(guidance["net_guidance_action"], float).tolist(), - net_guidance_norm=float(guidance["net_guidance_norm"]), - display_scale=float(guidance_scale), display_cap=float(guidance_cap), + components = ( + ("goal_guidance_action", CYAN, "goal"), + ("safety_guidance_action", MAGENTA, "safety"), ) + if all(key in guidance for key, _, _ in components): + component_metadata = {} + for key, color, label in components: + raw = np.asarray(guidance[key], float) + vector = float(guidance_scale) * raw + norm = float(np.linalg.norm(vector)) + if norm > float(guidance_cap): + vector *= float(guidance_cap) / norm + axis.add_patch(FancyArrowPatch( + tuple(start), tuple(start + vector), + arrowstyle="-|>", mutation_scale=15, + lw=2.6, color=color, shrinkA=0, shrinkB=0, zorder=11, + )) + component_metadata[label] = dict( + action=raw.tolist(), norm=float(np.linalg.norm(raw)), + ) + metadata.update( + guidance_present=True, + guidance_components=component_metadata, + guidance_semantics=guidance.get("component_semantics"), + display_scale=float(guidance_scale), + display_cap=float(guidance_cap), + ) + else: + vector = float(guidance_scale) * np.asarray( + guidance["net_guidance_action"], float, + ) + norm = float(np.linalg.norm(vector)) + if norm > float(guidance_cap): + vector *= float(guidance_cap) / norm + axis.add_patch(FancyArrowPatch( + tuple(start), tuple(start + vector), + arrowstyle="-|>", mutation_scale=15, + lw=2.6, color=MAGENTA, shrinkA=0, shrinkB=0, zorder=11, + )) + metadata.update( + guidance_present=True, + legacy_net_guidance=True, + net_guidance_action=np.asarray( + guidance["net_guidance_action"], float, + ).tolist(), + net_guidance_norm=float(guidance["net_guidance_norm"]), + display_scale=float(guidance_scale), + display_cap=float(guidance_cap), + ) else: metadata["guidance_present"] = False _set_clean_axis(axis) diff --git a/overnight_run_07_12_sfm/sfm_b1_ess_acquisition_viz.py b/overnight_run_07_12_sfm/sfm_b1_ess_acquisition_viz.py new file mode 100644 index 0000000..dfc3c7d --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_ess_acquisition_viz.py @@ -0,0 +1,201 @@ +"""Overlay actual executed D+ and D0 acquisition windows for ESS comparison.""" +from __future__ import annotations + +import argparse +from collections import defaultdict +import json +import math +import os + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import torch + +import _paths # noqa: F401 +import sfm_scene as SS + + +BLUE = "#0066ff" +MAGENTA = "#d000b5" +GAMMAS = (0.1, 0.3, 0.5, 1.0) + + +def _load(path): + return torch.load(os.path.abspath(path), map_location="cpu", weights_only=False) + + +def _rollout(state, controls): + state = np.asarray(state, np.float32).copy() + values = [state[:2].copy()] + for action in np.asarray(controls, np.float32): + state[:2] += SS.DT * state[2:4] + 0.5 * SS.DT ** 2 * action + state[2:4] += SS.DT * action + values.append(state[:2].copy()) + return np.asarray(values) + + +def _records(gather_dir): + executed = _load(os.path.join(gather_dir, "executed_round.pt")) + neutral = _load(os.path.join(gather_dir, "neutral_round.pt")) + contexts = {int(row["context_id"]): row for row in executed["contexts"]} + rows = [] + for window in executed["windows"]: + if int(window["y"]) != 1 or not bool(window["full_h"]): + continue + context = contexts[int(window["context_id"])] + rows.append({ + "population": "D+", + "scenario_id": int(context["scenario_id"]), + "gamma": float(context["gamma"]), + "step": int(context["step"]), + "state": np.asarray(context["state"], np.float32), + "controls": np.asarray(window["controls"], np.float32), + "sigma": float(window["sigma"]), + }) + for record in neutral["records"]: + if record["population"] != "D0" or int(record["verifier_y"]) != 0: + raise RuntimeError("neutral population semantics changed") + rows.append({ + "population": "D0", + "scenario_id": int(record["scenario_id"]), + "gamma": float(record["gamma"]), + "step": int(record["step"]), + "state": np.asarray(record["state"], np.float32), + "controls": np.asarray(record["controls"], np.float32), + "sigma": float(record["sigma"]), + }) + return rows + + +def _protocol(gather_dir): + payload = _load(os.path.join(gather_dir, "repair_trace.pt")) + protocol = payload["protocol"] + return { + "beta": float(protocol["beta"]), + "calibrated_ess_over_K": float(protocol["calibrated_ess_over_K"]), + "realized_ess_over_K": float(protocol["realized_ess_over_K"]), + "uplift": float(protocol["acquisition"]["uplift"]), + "gp_effective_cap": int(protocol["gp_selection"]["effective_cap"]), + } + + +def _diversity(rows): + if not rows: + return {"samples": 0, "occupied_025m_bins": 0, "heading_entropy": 0.0, + "spatial_rms": 0.0, "median_sigma": None} + xy = np.stack([row["state"][:2] for row in rows]) + bins = np.floor((xy - np.array([SS.TASK_LO, SS.TASK_LO])) / 0.25).astype(int) + occupied = len({tuple(value) for value in bins}) + actions = np.stack([row["controls"][0] for row in rows]) + angle = np.arctan2(actions[:, 1], actions[:, 0]) + hist, _ = np.histogram(angle, bins=12, range=(-math.pi, math.pi)) + probability = hist[hist > 0] / max(hist.sum(), 1) + entropy = float(-(probability * np.log(probability)).sum() / np.log(12.0)) + center = xy.mean(0) + return { + "samples": len(rows), + "occupied_025m_bins": occupied, + "heading_entropy": entropy, + "spatial_rms": float(np.sqrt(np.square(xy - center).sum(1).mean())), + "median_sigma": float(np.median([row["sigma"] for row in rows])), + } + + +def render(rows_by_name, labels, metadata, output_stem): + figure, axes = plt.subplots(3, 4, figsize=(13.0, 9.7), sharex=True, sharey=True) + manifest = {"gammas": list(GAMMAS), "protocol": metadata, "rows": {}} + for row_index, (name, rows) in enumerate(rows_by_name.items()): + manifest["rows"][name] = {} + for column, gamma in enumerate(GAMMAS): + axis = axes[row_index, column] + selected = [value for value in rows if abs(value["gamma"] - gamma) < 1e-6] + grouped = defaultdict(list) + for value in selected: + grouped[value["scenario_id"]].append(value) + for values in grouped.values(): + values.sort(key=lambda value: value["step"]) + trajectory = np.stack([value["state"][:2] for value in values]) + axis.plot(trajectory[:, 0], trajectory[:, 1], color="black", lw=0.55, alpha=0.5) + for value in selected: + branch = _rollout(value["state"], value["controls"]) + color = BLUE if value["population"] == "D+" else MAGENTA + axis.plot(branch[:, 0], branch[:, 1], color=color, lw=0.55, alpha=0.13) + for population, color, marker in (("D+", BLUE, "o"), ("D0", MAGENTA, "s")): + values = [value for value in selected if value["population"] == population] + if values: + xy = np.stack([value["state"][:2] for value in values]) + axis.scatter(xy[:, 0], xy[:, 1], s=8, marker=marker, color=color, + alpha=0.72, linewidths=0, zorder=3) + axis.scatter([0.0], [0.0], s=28, marker="o", facecolor="white", + edgecolor="black", linewidth=0.8, zorder=5) + axis.scatter([SS.GOAL[0]], [SS.GOAL[1]], s=75, marker="*", color="#f0b000", + edgecolor="black", linewidth=0.5, zorder=5) + axis.set_aspect("equal") + axis.set_xlim(SS.TASK_LO, SS.TASK_HI) + axis.set_ylim(SS.TASK_LO, SS.TASK_HI) + axis.grid(alpha=0.16, linewidth=0.5) + if row_index == 0: + axis.set_title(rf"$\gamma={gamma:g}$") + if column == 0: + axis.set_ylabel(labels[name]) + if row_index == 2: + axis.set_xlabel("x [m]") + positive = sum(value["population"] == "D+" for value in selected) + neutral = sum(value["population"] == "D0" for value in selected) + metrics = _diversity(selected) + metrics.update(Dplus=positive, D0=neutral) + manifest["rows"][name][str(gamma)] = metrics + axis.text(0.025, 0.975, + f"D+ {positive} D0 {neutral}\nbins {metrics['occupied_025m_bins']} " + f"Hθ {metrics['heading_entropy']:.2f}", + transform=axis.transAxes, va="top", ha="left", fontsize=7.5, + bbox=dict(facecolor="white", edgecolor="none", alpha=0.72, pad=1.5)) + handles = [ + plt.Line2D([0], [0], marker="o", linestyle="", color=BLUE, markersize=5, + label=r"executed exact-positive $D^+$"), + plt.Line2D([0], [0], marker="s", linestyle="", color=MAGENTA, markersize=5, + label=r"guided exact-negative neutral $D_0$"), + plt.Line2D([0], [0], color="black", lw=0.8, label="executed context path"), + ] + figure.legend(handles=handles, ncol=3, loc="upper center", frameon=False) + figure.suptitle("Actual D+ / D0 acquisition support · two shared OOD scenarios", y=0.965) + figure.tight_layout(rect=(0, 0, 1, 0.94)) + os.makedirs(os.path.dirname(os.path.abspath(output_stem)), exist_ok=True) + for suffix in ("png", "pdf"): + figure.savefig(f"{output_stem}.{suffix}", dpi=300, bbox_inches="tight") + plt.close(figure) + with open(f"{output_stem}.json", "w") as stream: + json.dump(manifest, stream, indent=2, sort_keys=True) + return manifest + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--pretrained-gather", required=True) + parser.add_argument("--ess05-gather", required=True) + parser.add_argument("--ess01-gather", required=True) + parser.add_argument("--output-stem", required=True) + args = parser.parse_args(argv) + rows = { + "pretrained": _records(args.pretrained_gather), + "ess05": _records(args.ess05_gather), + "ess01": _records(args.ess01_gather), + } + metadata = { + "pretrained": _protocol(args.pretrained_gather), + "ess05": _protocol(args.ess05_gather), + "ess01": _protocol(args.ess01_gather), + } + labels = { + "pretrained": "r0 gather\nGP empty", + "ess05": f"post-r1 · ESS 0.5\nβ={metadata['ess05']['beta']:.4g}", + "ess01": f"post-r1 · ESS 0.1\nβ={metadata['ess01']['beta']:.4g}", + } + render(rows, labels, metadata, os.path.abspath(args.output_stem)) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/sfm_b1_expand.py b/overnight_run_07_12_sfm/sfm_b1_expand.py index be7d6f7..c20c095 100644 --- a/overnight_run_07_12_sfm/sfm_b1_expand.py +++ b/overnight_run_07_12_sfm/sfm_b1_expand.py @@ -55,7 +55,10 @@ def validate(self): if (self.K, self.B, self.T, self.H, self.W, self.batch, self.lr, self.ess_target) != ( 16, 4, 180, 10, 2, 128, 1.0e-5, 0.5): raise ValueError("scientific B1 knobs differ from the frozen protocol") - if self.selector not in ("margin", "safemppi_cost"): + if self.selector not in ( + "margin", "progress_gated_margin", "safemppi_cost", + "balanced_rank", + ): raise ValueError("invalid arm selector") if self.scene_profile not in ( "legacy_velocity_ood", "requested_ood", "density_ood", diff --git a/overnight_run_07_12_sfm/sfm_b1_full_episode_audit.py b/overnight_run_07_12_sfm/sfm_b1_full_episode_audit.py new file mode 100644 index 0000000..3c21ff0 --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_full_episode_audit.py @@ -0,0 +1,553 @@ +"""Diagnostic-only full-episode B1 gathering with explicit post-NVP continuation. + +This module does not alter the fail-closed B1 trainer. It starts from the +pretrained policy, runs the ordinary K=16/B=4 RBF acquisition and the requested +execution selector, and records every resolved query. When the selected B +queries contain no admissible action, an independently sampled raw +temperature-one window is verified and its first action is executed so the +simulator can continue. + +That post-NVP transition is evidence gathering, not certified deployment. The +trace keeps the full-H verifier label, nominal-Hp gate, NVP event, progress/trap +event, and episode outcome separate. +""" +from __future__ import annotations + +import argparse +from collections import Counter, defaultdict +from concurrent.futures import ProcessPoolExecutor +import copy +import hashlib +import json +import os +import subprocess + +import numpy as np +import torch + +import _paths # noqa: F401 +import grid_policy_sfm as GPS +import sfm_b1_cost as BC +import sfm_b1_eval as BE +import sfm_b1_expand as BX +import sfm_b1_rbf as BR +import sfm_metrics2 as SM +import sfm_protocol as SP +import sfm_scene as SS + + +DEFAULT_SCENARIOS = (250_001, 250_003, 250_007) +DEFAULT_ELL = 0.24210826720721101 +DEFAULT_SAMPLE_SEED = 700_000 +DEFAULT_AUDIT_SEED = 20260723 +TRAP_HORIZON = 10 +TRAP_DISPLACEMENT = 0.2 + + +def _sha256_file(path): + digest = hashlib.sha256() + with open(path, "rb") as stream: + for chunk in iter(lambda: stream.read(1 << 20), b""): + digest.update(chunk) + return digest.hexdigest() + + +def _write_json(path, payload): + os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True) + temporary = os.fspath(path) + ".tmp" + with open(temporary, "w") as stream: + json.dump(payload, stream, indent=2, sort_keys=True) + os.replace(temporary, path) + + +def _save_torch(path, payload): + os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True) + temporary = os.fspath(path) + ".tmp" + torch.save(payload, temporary) + os.replace(temporary, path) + + +def _source(): + root = os.path.abspath(os.path.join(os.path.dirname(__file__), "..")) + commit = subprocess.check_output( + ["git", "rev-parse", "HEAD"], cwd=root, text=True + ).strip() + dirty = bool(subprocess.check_output( + ["git", "status", "--porcelain"], cwd=root, text=True + ).strip()) + return dict(commit=commit, tracked_worktree_clean=not dirty) + + +def _keyed_seed(base, *parts): + payload = json.dumps( + [int(base), *parts], separators=(",", ":"), sort_keys=False, + ).encode() + return int.from_bytes(hashlib.sha256(payload).digest()[:8], "little") % ( + 2 ** 63 - 1 + ) + + +@torch.no_grad() +def _keyed_windows( + policy, live, batch, *, K, round_i, step, source, seed, nfe, temp, +): + """Generate proposals and retain their exact Gaussian flow bases.""" + contexts = policy.ctx_from(batch["hp10"], batch["low"], batch["hist"]) + latent_parts = [] + for replica in live: + generator = np.random.default_rng(_keyed_seed( + seed, int(round_i), int(replica.scenario_id), + f"{float(replica.gamma):.8f}", int(step), str(source), + )) + latent_parts.append(generator.standard_normal( + (int(K), int(policy.d)), dtype=np.float32, + )) + x0 = torch.as_tensor( + np.stack(latent_parts), + device=contexts.device, + dtype=contexts.dtype, + ) + windows = BE.integrate_latents( + policy, + (x0 * float(temp)).reshape(-1, policy.d), + contexts.repeat_interleave(int(K), dim=0), + nfe=int(nfe), + ) + return ( + windows.reshape(len(live), int(K), int(policy.H_pred), 2), + contexts, + x0, + ) + + +@torch.no_grad() +def _features_from_x0(phi_policy, windows, contexts, x0, s): + K = int(windows.shape[1]) + features = phi_policy.phi_s_from_x0( + windows.reshape(-1, windows.shape[-2], 2), + contexts.repeat_interleave(K, dim=0), + x0.reshape(-1, phi_policy.d), + s=float(s), + ) + return BR.l2_normalize(features).reshape(len(contexts), K, -1) + + +@torch.no_grad() +def _calibrate_empty_gp_beta(phi_policy, gp, replicas, cfg, device): + live, batch = BX._stack_prepared(replicas, device) + windows, contexts, x0 = _keyed_windows( + phi_policy, live, batch, K=cfg.K, round_i=1, step=-1, + source="beta_calibration", seed=cfg.seed, + nfe=cfg.nfe, temp=cfg.temp, + ) + features = _features_from_x0( + phi_policy, windows, contexts, x0, cfg.phi_s, + ) + vectors = [] + for replica, values in zip(live, features): + generator = torch.Generator(device=values.device).manual_seed( + _keyed_seed( + cfg.seed, 1, replica.scenario_id, + f"{replica.gamma:.8f}", "beta_order", + ) + ) + order = torch.randperm( + len(values), generator=generator, device=values.device, + ) + vectors.extend(gp.sequential_score_vectors(values, order, cfg.B)) + return BR.solve_beta(vectors, target=cfg.ess_target) + + +def _trap(states, *, horizon=TRAP_HORIZON, displacement=TRAP_DISPLACEMENT): + if len(states) <= int(horizon): + return False + start = np.asarray(states[-int(horizon) - 1], float)[:2] + end = np.asarray(states[-1], float)[:2] + return bool(np.linalg.norm(end - start) < float(displacement)) + + +def _result_label(result): + if not result.get("resolved"): + return "verifier_error" + if int(result.get("y", 0)) == 0: + return "verifier_negative" + if not bool(result.get("full_h")) or int(result.get("terminal_step", -1)) != SP.H: + raise RuntimeError("the full-episode audit requires full-H=10 verifier semantics") + return "verifier_positive" + + +def _post_action_terminal(replica): + ped_xy, _ = SS.collect_humans(replica.humans) + clearance = float( + np.linalg.norm(ped_xy - replica.state[:2][None], axis=1).min() - SS.R_PED + ) + replica.minimum_clearance = min(replica.minimum_clearance, clearance) + collision = clearance < 0.0 + success = bool( + not collision and float(np.linalg.norm(replica.state[:2] - SS.GOAL)) < 0.5 + ) + if collision: + replica.alive = False + replica.status = "collision" + elif success: + replica.alive = False + replica.status = "success" + return collision, success, clearance + + +def collect( + checkpoint, *, scenarios=DEFAULT_SCENARIOS, gammas=SS.GAMMAS, + scene_profile="double_density_velocity_ood", device="cuda", + verifier_workers=32, sample_seed=DEFAULT_SAMPLE_SEED, + audit_seed=DEFAULT_AUDIT_SEED, ell=DEFAULT_ELL, T=SP.T, + selector="margin", outdir, +): + """Collect a fixed scenario-by-gamma full-episode diagnostic bundle.""" + scenarios = tuple(map(int, scenarios)) + gammas = tuple(map(float, gammas)) + if len(scenarios) != 3 or len(set(scenarios)) != 3: + raise ValueError("the requested audit requires exactly three distinct scenarios") + if gammas != tuple(map(float, SS.GAMMAS)): + raise ValueError(f"the requested audit requires all gammas={SS.GAMMAS}") + if scene_profile != "double_density_velocity_ood": + raise ValueError("this audit is pinned to the authenticated double-shift OOD") + if selector not in ("margin", "safemppi_cost", "balanced_rank"): + raise ValueError(f"unknown execution selector: {selector}") + if os.path.exists(outdir): + raise FileExistsError(f"refusing to reuse audit output: {outdir}") + + environment = SS.scene_profile(scene_profile) + policy, _ = GPS.load_sfm_policy(checkpoint, device=device) + policy.eval() + phi_policy = copy.deepcopy(policy).eval() + for parameter in phi_policy.parameters(): + parameter.requires_grad_(False) + + replicas = [ + BX.Replica( + scenario, gamma, n_ped=environment["n_ped"], + ped_speed_range=tuple(environment["ped_speed_range"]), + ) + for scenario in scenarios for gamma in gammas + ] + cfg = BX.ArmConfig( + name="diagnostic", selector=selector, alpha=0.0, rounds=1, + scene_profile=scene_profile, verifier_workers=int(verifier_workers), + seed=int(audit_seed), + ).validate() + gp = BR.RBFGP(float(ell), cfg.gp_lam) + beta, calibrated_ess = _calibrate_empty_gp_beta( + phi_policy, gp, replicas, cfg, device, + ) + + traces = [] + counts = Counter() + ess_values = [] + sigma_pool, sigma_selected = [], [] + trap_active = defaultdict(bool) + with ProcessPoolExecutor(max_workers=int(verifier_workers)) as executor: + for step in range(int(T)): + live = [replica for replica in replicas if replica.alive] + live, batch = BX._stack_prepared(live, device) + if not live: + break + with torch.no_grad(): + audit_windows, contexts, x0 = _keyed_windows( + policy, live, batch, K=cfg.K, round_i=1, step=step, + source="K", seed=int(sample_seed), + nfe=cfg.nfe, temp=cfg.temp, + ) + raw_windows, _, raw_x0 = _keyed_windows( + policy, live, batch, K=1, round_i=1, step=step, + source="raw_continuation", seed=int(sample_seed), + nfe=cfg.nfe, temp=cfg.temp, + ) + raw_windows = raw_windows[:, 0] + raw_x0 = raw_x0[:, 0] + features = _features_from_x0( + phi_policy, audit_windows, contexts, x0, cfg.phi_s, + ) + + selected_by_context = [] + acquisition_by_context = [] + for context_index, replica in enumerate(live): + acquisition_generator = torch.Generator( + device=features.device, + ).manual_seed(_keyed_seed( + int(audit_seed), 1, replica.scenario_id, + f"{replica.gamma:.8f}", step, "acquisition", + )) + selected, acquisition = gp.sequential_acquire( + features[context_index], cfg.B, beta, + generator=acquisition_generator, + ) + selected_by_context.append(selected) + acquisition_by_context.append(acquisition) + sigma_pool.extend(map( + float, gp.acquisition_sigma(features[context_index]).detach().cpu() + )) + sigma_selected.extend(float(row["chosen_sigma"]) for row in acquisition) + ess_values.extend(float(row["ess_norm"]) for row in acquisition) + + tasks = [] + for context_index, replica in enumerate(live): + prepared = replica.prepared + for candidate_id in selected_by_context[context_index]: + tasks.append(( + context_index, candidate_id, prepared["state"], + audit_windows[context_index, candidate_id].detach().cpu().numpy(), + prepared["ped_xy"], prepared["ped_vel"], replica.gamma, + )) + results = list(executor.map(SM.verify_in_worker, tasks)) + by_context = defaultdict(dict) + for context_index, candidate_id, result in results: + by_context[int(context_index)][int(candidate_id)] = result + + prepared_contexts = [] + raw_tasks = [] + for context_index, replica in enumerate(live): + prepared = replica.prepared + pedestrian_prediction = SM.predict_pedestrians( + prepared["ped_xy"], prepared["ped_vel"], cfg.H, + ) + all_rows = [] + for candidate_id in range(cfg.K): + controls = audit_windows[ + context_index, candidate_id + ].detach().cpu().numpy() + segment = SM.rollout_positions(prepared["state"], controls) + all_rows.append(dict( + candidate_id=candidate_id, controls=controls, segment=segment, + x0=x0[context_index, candidate_id].detach().cpu().numpy(), + mode=BE.classify_candidate(segment, pedestrian_prediction), + )) + query_rows = [] + for acquisition_step, candidate_id in enumerate( + selected_by_context[context_index]): + result = by_context[context_index][candidate_id] + query_rows.append(dict( + candidate_id=int(candidate_id), + controls=all_rows[candidate_id]["controls"], + result=result, mode=all_rows[candidate_id]["mode"], + acquisition_step=int(acquisition_step), + sigma=float( + acquisition_by_context[context_index][ + acquisition_step + ]["chosen_sigma"] + ), + )) + counts[f"B_{_result_label(result)}"] += 1 + chosen = BC.select_admissible( + query_rows, selector=selector, state=prepared["state"], + ped_xy=prepared["ped_xy"], ped_vel=prepared["ped_vel"], + gamma=replica.gamma, + ) + prepared_contexts.append((all_rows, query_rows, chosen)) + if chosen is None: + raw_tasks.append(( + context_index, -1, prepared["state"], + raw_windows[context_index].detach().cpu().numpy(), + prepared["ped_xy"], prepared["ped_vel"], replica.gamma, + )) + + for context_index, candidate_id, result in executor.map( + SM.verify_in_worker, raw_tasks): + by_context[int(context_index)][int(candidate_id)] = result + + for context_index, replica in enumerate(live): + prepared = replica.prepared + all_rows, query_rows, chosen = prepared_contexts[context_index] + + raw_controls = raw_windows[context_index].detach().cpu().numpy() + raw_base = raw_x0[context_index].detach().cpu().numpy() + nvp_context = chosen is None + if chosen is None: + raw_result = by_context[context_index][-1] + raw_margin, raw_hp_old, raw_hp_new = BC.nominal_hp_margin( + prepared["state"], raw_controls[0], prepared["ped_xy"], + replica.gamma, + ) + raw_admissible = bool( + raw_result.get("resolved") + and int(raw_result.get("y", 0)) == 1 + and bool(raw_result.get("full_h")) + and raw_margin >= -1.0e-9 + ) + executed_controls = raw_controls + executed_result = raw_result + executed_id = None + execution_source = ( + "certified_raw_rescue" if raw_admissible + else "uncertified_raw_continuation" + ) + counts["B_NVP_context"] += 1 + counts[f"raw_continuation_{_result_label(raw_result)}"] += 1 + raw_candidate = dict( + controls=np.asarray(raw_controls, np.float32), + x0=np.asarray(raw_base, np.float32), + result=raw_result, hp_margin=float(raw_margin), + hp_old=float(raw_hp_old), hp_new=float(raw_hp_new), + admissible=raw_admissible, + ) + else: + executed_controls = chosen["controls"] + executed_x0 = all_rows[int(chosen["candidate_id"])]["x0"] + executed_result = chosen["result"] + executed_id = int(chosen["candidate_id"]) + execution_source = f"verified_{selector}" + raw_candidate = None + if chosen is None: + executed_x0 = raw_base + executed_label = _result_label(executed_result) + counts[f"executed_{executed_label}"] += 1 + counts[f"source_{execution_source}"] += 1 + + before = prepared["state"].copy() + BX._advance(replica, executed_controls[0]) + trap_event = _trap(replica.states) + trap_key = (replica.scenario_id, replica.gamma) + trap_entry = bool(trap_event and not trap_active[trap_key]) + trap_active[trap_key] = bool(trap_event) + collision_after, success_after, clearance_after = _post_action_terminal( + replica + ) + if trap_entry: + counts["trap_entries"] += 1 + if collision_after: + counts["collision_events"] += 1 + if success_after: + counts["success_events"] += 1 + + negative_reasons = [] + if nvp_context: + negative_reasons.append("finite_B_NVP") + if executed_label == "verifier_negative": + negative_reasons.append("executed_full_H_rejected") + elif executed_label == "verifier_error": + negative_reasons.append("executed_verifier_error") + if ( + executed_label == "verifier_positive" + and not execution_source.startswith("verified_") + ): + if raw_margin < -1.0e-9: + negative_reasons.append("executed_nominal_Hp_gate_failure") + if trap_event: + negative_reasons.append("ten_step_progress_below_0p2m") + if collision_after: + negative_reasons.append("collision") + + traces.append(dict( + round=1, step=int(step), scenario_id=replica.scenario_id, + gamma=replica.gamma, state=before, + next_state=replica.state.copy(), + ped_xy=prepared["ped_xy"], ped_vel=prepared["ped_vel"], + all_K=all_rows, + selected_ids=list(map(int, selected_by_context[context_index])), + query_rows=query_rows, acquisition=acquisition_by_context[context_index], + executed_id=executed_id, + executed_controls=np.asarray(executed_controls, np.float32), + executed_x0=np.asarray(executed_x0, np.float32), + executed_result=executed_result, + executed_label=executed_label, + execution_source=execution_source, + nvp_context=bool(nvp_context), + raw_candidate=raw_candidate, + trap_event=bool(trap_event), + trap_entry=bool(trap_entry), + collision_after_action=bool(collision_after), + success_after_action=bool(success_after), + clearance_after_action=float(clearance_after), + negative_reasons=negative_reasons, + )) + + for replica in replicas: + if replica.alive: + replica.alive = False + replica.status = "timeout" + outcomes = [dict( + scenario_id=replica.scenario_id, gamma=replica.gamma, + status=replica.status, steps=len(replica.controls), + success=replica.status == "success", + collision=replica.status == "collision", + timeout=replica.status == "timeout", + minimum_clearance=float(replica.minimum_clearance), + ) for replica in replicas] + counts.update(f"outcome_{row['status']}" for row in outcomes) + + source = _source() + bundle = dict( + version=2, status="SFM_B1_FULL_EPISODE_LABEL_AUDIT_COMPLETE", + diagnostic_only=True, enters_training_or_gp=False, + certified_deployment=False, + continuation_semantics=( + f"verified {selector} B action when available; otherwise independently " + "sampled raw temp=1 action is executed after being labeled, even when " + "uncertified, solely to continue the offline simulator diagnostic" + ), + label_semantics=dict( + safety="executed_label is the independent exact full-H=10 verifier result", + nvp="finite B=4 acquisition event; not itself a verifier-negative label", + trap=( + f"separate performance event: displacement over {TRAP_HORIZON} " + f"executed actions is below {TRAP_DISPLACEMENT} m" + ), + collision="episode event; never retroactively relabels earlier windows", + ), + source=source, checkpoint=os.path.abspath(checkpoint), + checkpoint_sha256=_sha256_file(checkpoint), + environment=environment, scenarios=list(scenarios), gammas=list(gammas), + sample_seed=int(sample_seed), audit_seed=int(audit_seed), + protocol=dict( + K=cfg.K, B=cfg.B, H=cfg.H, T=int(T), selector=selector, + ell=float(ell), gp_buffer=0, beta=float(beta), + representation=( + "normalize(phi_theta((1-s)*x0+s*U/u_max,s,c)); " + "stored proposal-specific x0; s=0.9" + ), + calibrated_ess_over_K=float(calibrated_ess), + realized_ess_over_K=float(np.mean(ess_values)), + acquisition=BR.acquisition_diagnostics(sigma_pool, sigma_selected), + ), + counts=dict(counts), outcomes=outcomes, traces=traces, + ) + os.makedirs(outdir) + trace_path = os.path.join(outdir, "full_episode_label_audit.pt") + _save_torch(trace_path, bundle) + manifest = { + key: value for key, value in bundle.items() if key != "traces" + } + manifest["trace_path"] = os.path.abspath(trace_path) + manifest["trace_sha256"] = _sha256_file(trace_path) + _write_json(os.path.join(outdir, "full_episode_label_audit.json"), manifest) + return manifest + + +def main(argv=None): + parser = argparse.ArgumentParser() + parser.add_argument("--checkpoint", required=True) + parser.add_argument("--outdir", required=True) + parser.add_argument("--scenarios", nargs=3, type=int, default=DEFAULT_SCENARIOS) + parser.add_argument("--scene-profile", default="double_density_velocity_ood") + parser.add_argument("--device", default="cuda") + parser.add_argument("--verifier-workers", type=int, default=32) + parser.add_argument("--sample-seed", type=int, default=DEFAULT_SAMPLE_SEED) + parser.add_argument("--audit-seed", type=int, default=DEFAULT_AUDIT_SEED) + parser.add_argument("--ell", type=float, default=DEFAULT_ELL) + parser.add_argument("--T", type=int, default=SP.T) + parser.add_argument( + "--selector", + choices=("margin", "safemppi_cost", "balanced_rank"), + default="margin", + ) + args = parser.parse_args(argv) + collect( + args.checkpoint, scenarios=args.scenarios, + scene_profile=args.scene_profile, device=args.device, + verifier_workers=args.verifier_workers, + sample_seed=args.sample_seed, audit_seed=args.audit_seed, + ell=args.ell, T=args.T, selector=args.selector, outdir=args.outdir, + ) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/sfm_b1_full_episode_viz.py b/overnight_run_07_12_sfm/sfm_b1_full_episode_viz.py new file mode 100644 index 0000000..825a702 --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_full_episode_viz.py @@ -0,0 +1,290 @@ +"""Render the diagnostic full-episode B1 label audit. + +Blue/red trajectory segments are exact full-H verifier labels of the action +window that was actually executed. A red NVP ring is a separate finite-B +context event. Green H=10 levels are drawn only for an executed verifier +positive; post-NVP uncertified raw continuations never acquire green geometry. +""" +from __future__ import annotations + +import argparse +from collections import defaultdict +import hashlib +import json +import os + +import matplotlib +matplotlib.use("Agg") +import matplotlib.animation as animation +import matplotlib.pyplot as plt +from matplotlib.lines import Line2D +import numpy as np +import torch + +import _paths # noqa: F401 +import sfm_b1_density_viz as DV +import sfm_b1_viz as BV +import sfm_scene as SS + + +def _sha256(path): + digest = hashlib.sha256() + with open(path, "rb") as stream: + for chunk in iter(lambda: stream.read(1 << 20), b""): + digest.update(chunk) + return digest.hexdigest() + + +def _write_json(path, payload): + os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True) + temporary = os.fspath(path) + ".tmp" + with open(temporary, "w") as stream: + json.dump(payload, stream, indent=2, sort_keys=True) + os.replace(temporary, path) + + +def _index(traces): + output = defaultdict(dict) + for trace in traces: + key = (int(trace["scenario_id"]), round(float(trace["gamma"]), 8)) + step = int(trace["step"]) + if step in output[key]: + raise ValueError(f"duplicate trace key {key + (step,)}") + output[key][step] = trace + return output + + +def _executed_color(trace): + if trace["executed_label"] == "verifier_positive": + return BV.BLUE + if trace["executed_label"] == "verifier_negative": + return BV.RED + return BV.GRAY + + +def _draw_history(axis, rows, step): + for index in sorted(value for value in rows if value <= int(step)): + trace = rows[index] + before = np.asarray(trace["state"], float)[:2] + after = np.asarray(trace["next_state"], float)[:2] + color = _executed_color(trace) + axis.plot( + [before[0], after[0]], [before[1], after[1]], + color=color, lw=2.2, marker=".", ms=2.0, alpha=.94, zorder=6, + ) + + +def _draw_candidates(axis, trace): + selected = set(map(int, trace["selected_ids"])) + for row in trace["all_K"]: + path = np.asarray(row["segment"], float) + axis.plot( + path[:, 0], path[:, 1], color=BV.GRAY, lw=.38, + marker=".", ms=1.0, alpha=.22, zorder=3, + ) + for candidate_id in sorted(selected): + row = BV._trace_candidate(trace, candidate_id) + path = np.asarray(row["segment"], float) + axis.plot( + path[:, 0], path[:, 1], color=BV.ORANGE, lw=.72, + marker=".", ms=1.3, alpha=.92, zorder=4, + ) + for candidate_id in sorted(selected): + status, query = BV._candidate_status(trace, candidate_id) + if status not in ("positive", "negative"): + continue + path = np.asarray(BV._trace_candidate(trace, candidate_id)["segment"], float) + color = BV.GREEN if status == "positive" else BV.RED + axis.plot( + path[:, 0], path[:, 1], color=color, lw=1.0, + marker=".", ms=1.45, alpha=.96, zorder=5, + ) + if status == "negative": + axis.plot( + path[-1, 0], path[-1, 1], "x", color=BV.RED, + ms=3.5, mew=.9, zorder=8, + ) + + +def _draw_executed(axis, trace): + result = trace["executed_result"] + path = np.asarray(result.get( + "segment", + trace["executed_controls"], + ), float) + if path.shape != (11, 2): + path = np.asarray([ + np.asarray(trace["state"], float)[:2], + np.asarray(trace["next_state"], float)[:2], + ]) + color = _executed_color(trace) + axis.plot(path[:, 0], path[:, 1], color=color, lw=.75, alpha=.62, zorder=6) + axis.plot(path[:2, 0], path[:2, 1], color=color, lw=3.0, zorder=9) + axis.annotate( + "", xy=path[1], xytext=path[0], + arrowprops=dict(arrowstyle="->", color=color, lw=2.3), + ) + if trace["executed_label"] == "verifier_positive": + query = dict(result=result) + audit = DV.checked_verifier_levels(trace, query, H=10) + DV._draw_verifier_geometry(axis, audit) + + +def draw_cell(axis, rows, step): + available = [value for value in rows if value <= int(step)] + current_step = max(available) if available else min(rows) + trace = rows[current_step] + BV._draw_common(axis, trace, nominal_levels=False) + _draw_history(axis, rows, current_step) + _draw_candidates(axis, trace) + _draw_executed(axis, trace) + position = np.asarray(trace["state"], float)[:2] + if trace["nvp_context"]: + axis.plot( + position[0], position[1], marker="o", ms=10, mfc="none", + mec=BV.RED, mew=1.6, zorder=12, + ) + if trace["trap_entry"]: + axis.plot(position[0], position[1], marker="s", ms=6, + mfc="none", mec=BV.RED, mew=1.2, zorder=12) + if trace["collision_after_action"]: + after = np.asarray(trace["next_state"], float)[:2] + axis.plot(after[0], after[1], marker="x", ms=8, + color=BV.RED, mew=1.8, zorder=13) + DV._set_clean_axis(axis) + return trace + + +def _legend(): + return [ + Line2D([], [], color=BV.GRAY, lw=.7, label="K=16 generated"), + Line2D([], [], color=BV.ORANGE, lw=1.1, label="B=4 RBF queried"), + Line2D([], [], color=BV.GREEN, lw=1.4, label="B full-H positive"), + Line2D([], [], color=BV.RED, lw=1.4, marker="x", label="B full-H rejected"), + Line2D([], [], color=BV.BLUE, lw=2.7, label="executed window: full-H positive"), + Line2D([], [], color=BV.RED, lw=2.7, label="executed window: full-H rejected"), + Line2D([], [], color=BV.GREEN, lw=.7, label="executed verifier levels h=1..10"), + Line2D([], [], marker="o", ms=8, mfc="none", mec=BV.RED, lw=0, + label="finite-B NVP context"), + Line2D([], [], marker="s", ms=6, mfc="none", mec=BV.RED, lw=0, + label="first entry: 10-step progress < 0.2 m"), + ] + + +def render(trace_path, output_mp4, output_png, output_json, *, fps=5, frame_stride=2): + if int(fps) <= 0 or int(frame_stride) <= 0: + raise ValueError("fps and frame_stride must be positive") + bundle = torch.load(trace_path, map_location="cpu", weights_only=False) + if bundle.get("status") != "SFM_B1_FULL_EPISODE_LABEL_AUDIT_COMPLETE": + raise ValueError("input is not a completed full-episode audit") + scenarios = tuple(map(int, bundle["scenarios"])) + gammas = tuple(map(float, bundle["gammas"])) + if len(scenarios) != 3 or gammas != tuple(map(float, SS.GAMMAS)): + raise ValueError("renderer requires three scenarios and all seven gammas") + index = _index(bundle["traces"]) + missing = [ + (scenario, gamma) for scenario in scenarios for gamma in gammas + if (scenario, round(gamma, 8)) not in index + ] + if missing: + raise ValueError(f"missing audit cells: {missing}") + + maximum = max(max(rows) for rows in index.values()) + frames = list(range(0, maximum + 1, int(frame_stride))) + if frames[-1] != maximum: + frames.append(maximum) + figure, axes = plt.subplots(3, 7, figsize=(23.2, 10.1)) + figure.subplots_adjust( + left=.035, right=.815, bottom=.025, top=.94, wspace=.025, hspace=.04, + ) + for column, gamma in enumerate(gammas): + figure.text( + .035 + (.78 / 7) * (column + .5), .965, f"$\\gamma={gamma:g}$", + ha="center", va="center", fontsize=10, + ) + for row, scenario in enumerate(scenarios): + figure.text( + .012, .94 - (.915 / 3) * (row + .5), f"episode\n{scenario}", + ha="center", va="center", rotation=90, fontsize=9, + ) + figure.legend( + handles=_legend(), loc="center left", bbox_to_anchor=(.825, .58), + frameon=False, fontsize=8, + ) + figure.text( + .825, .25, + "Offline diagnostic only\n" + "NVP does not stop this simulator trace.\n" + "After NVP, a separately sampled raw\n" + "temperature-1 first action advances it.\n" + "Red continuation is not certified safety.", + ha="left", va="top", fontsize=8, + ) + + final_cells = {} + + def update(step): + final_cells.clear() + for row, scenario in enumerate(scenarios): + for column, gamma in enumerate(gammas): + axis = axes[row, column] + axis.clear() + trace = draw_cell( + axis, index[(scenario, round(gamma, 8))], int(step) + ) + final_cells[f"{scenario}:{gamma:g}"] = dict( + rendered_step=int(trace["step"]), + execution_source=trace["execution_source"], + executed_label=trace["executed_label"], + nvp_context=bool(trace["nvp_context"]), + ) + return [] + + for path in (output_mp4, output_png, output_json): + os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True) + movie = animation.FuncAnimation( + figure, update, frames=frames, interval=1000 / int(fps), blit=False, + ) + movie.save( + output_mp4, writer=animation.FFMpegWriter(fps=int(fps), bitrate=4200), + dpi=105, + ) + update(maximum) + figure.savefig(output_png, dpi=165, bbox_inches="tight") + plt.close(figure) + + report = dict( + status="SFM_B1_FULL_EPISODE_LABEL_VIZ_COMPLETE", + diagnostic_only=True, trace_path=os.path.abspath(trace_path), + trace_sha256=_sha256(trace_path), scenarios=list(scenarios), + gammas=list(gammas), frame_stride=int(frame_stride), fps=int(fps), + frames=frames, mp4=os.path.abspath(output_mp4), + mp4_sha256=_sha256(output_mp4), png=os.path.abspath(output_png), + png_sha256=_sha256(output_png), + color_semantics=( + "blue/red trail is exact full-H label of the actually executed " + "window; NVP/trap/collision remain separate event markers" + ), + final_cells=dict(final_cells), + ) + _write_json(output_json, report) + return report + + +def main(argv=None): + parser = argparse.ArgumentParser() + parser.add_argument("--trace", required=True) + parser.add_argument("--output-mp4", required=True) + parser.add_argument("--output-png", required=True) + parser.add_argument("--output-json", required=True) + parser.add_argument("--fps", type=int, default=5) + parser.add_argument("--frame-stride", type=int, default=2) + args = parser.parse_args(argv) + render( + args.trace, args.output_mp4, args.output_png, args.output_json, + fps=args.fps, frame_stride=args.frame_stride, + ) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/sfm_b1_kazuki_repair.py b/overnight_run_07_12_sfm/sfm_b1_kazuki_repair.py new file mode 100644 index 0000000..70c57a0 --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_kazuki_repair.py @@ -0,0 +1,193 @@ +"""Same-latent locked-Kazuki guidance for B1 acquisition repair. + +This module deliberately implements only the locked Kazuki *guidance field*. +It does not add Kazuki's 200-sample generator, warm start, MPPI refinement, +output filter, privileged SFM lookahead, or any new proposal template. + +For each selected B candidate, the original Gaussian base ``x0`` is reused and +integrated through the current policy with the fixed goal/CBF coefficients. +The result is therefore a guided regeneration of the same latent lineage, not +an action-space edit of the already generated control window. +""" +from __future__ import annotations + +from dataclasses import replace + +import numpy as np +import torch + +import sfm_b1_full_episode_audit as FA +import sfm_kazuki as KZ +import sfm_metrics2 as SM +import sfm_scene as SS + + +LOCKED_SAFE_COEF = 0.3 +LOCKED_GOAL_COEF = 0.5 +TRAP_HORIZON = 10 +TRAP_DISPLACEMENT = 0.2 +TRAP_PATIENCE = 3 + + +def locked_guidance_config(gamma, *, n_sample): + """Materialize the immutable Kazuki guidance coefficients for one B pool.""" + base = KZ.KazukiConfig( + safe_coefs=(LOCKED_SAFE_COEF,), + goal_coef=LOCKED_GOAL_COEF, + safe_coef_gamma_span=0.0, + goal_coef_gamma_span=0.0, + ).validate() + controller = KZ._gamma_controller_config(base, float(gamma)).validate() + guidance = KZ._gamma_guidance_config(controller, float(gamma)) + return replace(guidance, n_sample=int(n_sample)).validate() + + +def collector_ode_times(nfe): + """The same Euler knot schedule used by the unguided B1 K proposals.""" + nfe = int(nfe) + if nfe < 1: + raise ValueError("nfe must be positive") + return tuple(index / nfe for index in range(nfe + 1)) + + +def same_latent_guided_controls( + policy, + context, + state, + ped_xy, + ped_vel, + gamma, + x0, + *, + nfe=8, + collect_diagnostics=True, +): + """Regenerate selected B latent lineages under fixed Kazuki guidance.""" + x0 = torch.as_tensor( + x0, device=context.device, dtype=context.dtype, + ) + if x0.ndim != 2 or x0.shape[1] != int(policy.d): + raise ValueError(f"x0 must have shape [B,{policy.d}], got {tuple(x0.shape)}") + count = int(len(x0)) + if count < 1: + raise ValueError("at least one selected latent is required") + cfg = locked_guidance_config(gamma, n_sample=count) + goal = torch.as_tensor( + SS.GOAL, device=context.device, dtype=context.dtype, + ) + ped_prediction = KZ.predict_pedestrians_t( + ped_xy, ped_vel, int(policy.H_pred), SS.DT, + context.device, context.dtype, + ) + ped_velocity = torch.as_tensor( + ped_vel, device=context.device, dtype=context.dtype, + ) + guided, ode_trace, unguided = KZ.guided_generate( + policy, + context, + np.asarray(state, np.float32), + goal, + ped_prediction, + ped_velocity, + SS.R_PED + float(cfg.collision_margin), + x0, + collector_ode_times(nfe), + cfg, + collect_diagnostics=bool(collect_diagnostics), + ) + guided_controls = torch.clamp( + guided.reshape(count, policy.H_pred, 2) * float(policy.u_max), + -float(policy.u_max), + float(policy.u_max), + ) + diagnostics = dict( + operator="same_latent_locked_kazuki_guidance", + candidate_count=count, + safe_coef=LOCKED_SAFE_COEF, + goal_coef=LOCKED_GOAL_COEF, + nfe=int(nfe), + ode_times=collector_ode_times(nfe), + no_mppi_refinement=True, + no_new_latents=True, + ) + if collect_diagnostics: + unguided_controls = torch.clamp( + unguided.reshape(count, policy.H_pred, 2) * float(policy.u_max), + -float(policy.u_max), + float(policy.u_max), + ) + terminal = ode_trace[-1] + diagnostics.update( + unguided_controls=unguided_controls.detach().cpu().numpy(), + net_first_action=( + guided_controls[:, 0] - unguided_controls[:, 0] + ).detach().cpu().numpy(), + goal_first_action=( + terminal["integrated_goal_guidance"] + .reshape(count, policy.H_pred, 2)[:, 0] + * float(policy.u_max) + ).detach().cpu().numpy(), + safety_first_action=( + terminal["integrated_safety_guidance"] + .reshape(count, policy.H_pred, 2)[:, 0] + * float(policy.u_max) + ).detach().cpu().numpy(), + component_semantics=terminal["component_semantics"], + ) + return guided_controls, diagnostics + + +def predicted_trap(states, first_action, *, horizon=TRAP_HORIZON, + displacement=TRAP_DISPLACEMENT): + """Whether executing ``first_action`` makes the declared trap predicate true.""" + states = [np.asarray(value, np.float32) for value in states] + if len(states) < int(horizon): + return False + current = states[-1] + action = np.asarray(first_action, np.float32) + next_position = ( + current[:2] + SS.DT * current[2:4] + + 0.5 * SS.DT ** 2 * action + ) + return bool( + np.linalg.norm(next_position - states[-int(horizon)][:2]) + < float(displacement) + ) + + +def post_action_trap_matches(states, first_action): + """Audit that the predictive predicate equals the existing post-action test.""" + states = [np.asarray(value, np.float32) for value in states] + current = states[-1] + action = np.asarray(first_action, np.float32) + next_position = SM.rollout_positions(current, action[None])[-1] + next_state = np.concatenate([ + next_position, + current[2:4] + SS.DT * action, + ]).astype(np.float32) + expected = predicted_trap(states, action) + observed = FA._trap(states + [next_state]) + if bool(expected) != bool(observed): + raise AssertionError("predictive and post-action trap predicates disagree") + return bool(expected) + + +def next_trap_streak(current_streak, trap_event): + """Update the per-lineage streak and report the fail-closed boundary.""" + streak = int(current_streak) + 1 if bool(trap_event) else 0 + return streak, bool(streak >= TRAP_PATIENCE) + + +def manifest(): + return dict( + operator="same-latent B4 Kazuki guidance repair", + safe_coef=LOCKED_SAFE_COEF, + goal_coef=LOCKED_GOAL_COEF, + trap_horizon=TRAP_HORIZON, + trap_displacement=TRAP_DISPLACEMENT, + trap_patience=TRAP_PATIENCE, + exclusions=( + "no privileged MPC; no independent raw fallback; no new latent; " + "no Kazuki MPPI refinement; no output shield" + ), + ) diff --git a/overnight_run_07_12_sfm/sfm_b1_kazuki_repair_audit.py b/overnight_run_07_12_sfm/sfm_b1_kazuki_repair_audit.py new file mode 100644 index 0000000..8057345 --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_kazuki_repair_audit.py @@ -0,0 +1,1199 @@ +"""Training-ready audit of same-latent Kazuki repair during B1 gathering. + +The collector keeps three deliberately separate stores: + +* ``executed_round.pt``: one exact full-H window per executed context, matching + the existing offline replay contract; +* ``query_sidecar.pt``: every resolved base B=4 and repair B=4 query, retained + for audit only unless a later experiment explicitly opts into all-query + replay. +* ``neutral_round.pt``: opt-in guided verifier-negative executions, isolated + from training D+/D-, GP, and replay while remaining exact-negative in the + audit query sidecar. + +No independent raw continuation and no privileged MPC proposal are used. +""" +from __future__ import annotations + +import argparse +from collections import Counter, defaultdict +from concurrent.futures import ProcessPoolExecutor +from contextlib import nullcontext +import copy +from dataclasses import replace +import os + +import numpy as np +import torch + +import _paths # noqa: F401 +import grid_policy_sfm as GPS +import sfm_b1_cost as BC +import sfm_b1_eval as BE +import sfm_b1_expand as BX +import sfm_b1_full_episode_audit as FA +import sfm_b1_kazuki_repair as KR +import sfm_b1_offline_store as OS +import sfm_b1_rbf as BR +import sfm_b1_store as BS +import sfm_metrics2 as SM +import sfm_scene as SS + + +STATUS = "SFM_B1_KAZUKI_REPAIR_AUDIT_COMPLETE" +DEFAULT_SCENARIOS = (250_001, 250_003) +DEFAULT_GAMMAS = (0.1, 0.5, 1.0) +DEFAULT_ELL = 0.24210826720721101 +DEFAULT_SAMPLE_SEED = 700_000 +DEFAULT_AUDIT_SEED = 20260730 +EXPECTED_CHECKPOINT_SHA256 = ( + "1b5179c935d3eeff8824967d707d64cc9bab273949ee1f0e4f190172bab1b215" +) + + +def _query_row( + candidate_id, controls, x0, result, *, acquisition_step, sigma, + mode, source, parent_candidate_id=None, +): + return dict( + candidate_id=int(candidate_id), + parent_candidate_id=( + None if parent_candidate_id is None else int(parent_candidate_id) + ), + controls=np.asarray(controls, np.float32), + x0=np.asarray(x0, np.float32), + result=result, + acquisition_step=int(acquisition_step), + sigma=float(sigma), + mode=str(mode), + source=str(source), + ) + + +def _add_sidecar_query(shard, context_id, row): + result = row["result"] + if not result.get("resolved"): + shard.add_error( + context_key=( + shard.contexts[int(context_id)]["scenario_id"], + shard.contexts[int(context_id)]["gamma"], + shard.contexts[int(context_id)]["step"], + ), + candidate_id=row["candidate_id"], + error=result.get("error"), + ) + return None + query_id = shard.add_resolved_query( + context_id, + row["candidate_id"], + row["controls"], + row["sigma"], + result, + acquisition_step=row["acquisition_step"], + hp_margin=row.get("hp_margin"), + expert_cost=row.get("expert_cost"), + mode=row["mode"], + ) + stored = shard.queries[int(query_id)] + stored.update( + x0=np.asarray(row["x0"], np.float32), + query_source=row["source"], + parent_candidate_id=row["parent_candidate_id"], + audit_only=True, + replay_default=False, + ) + row["query_id"] = int(query_id) + return int(query_id) + + +def _admissible(rows, selector, prepared, gamma): + return BC.select_admissible( + [row for row in rows if row["result"].get("resolved")], + selector=selector, + state=prepared["state"], + ped_xy=prepared["ped_xy"], + ped_vel=prepared["ped_vel"], + gamma=float(gamma), + ) + + +def _all_full_h_negative(rows, expected=4): + return len(rows) == int(expected) and all( + row["result"].get("resolved") + and int(row["result"].get("y", -1)) == 0 + and bool(row["result"].get("full_h")) + and int(row["result"].get("terminal_step", -1)) == 10 + for row in rows + ) + + +def _select_neutral(rows, selector, prepared, gamma): + """Rank four exact-negative guided rows without relabeling them.""" + if not _all_full_h_negative(rows): + return None + for row in rows: + margin, hp_old, hp_new = BC.nominal_hp_margin( + prepared["state"], + row["controls"][0], + prepared["ped_xy"], + gamma, + ) + row["hp_margin"] = float(margin) + row["hp_old"] = float(hp_old) + row["hp_new"] = float(hp_new) + if selector == "margin": + return max( + rows, + key=lambda row: ( + row["hp_margin"], + -int(row["candidate_id"]), + ), + ) + if selector == "progress_gated_margin": + return BC.select_progress_gated_margin( + rows, state=prepared["state"], + ) + if selector != "safemppi_cost": + raise ValueError(f"unknown neutral selector: {selector}") + controls = torch.as_tensor( + np.stack([row["controls"] for row in rows]), + dtype=torch.float32, + ) + costs = BC.safemppi_proposal_cost( + prepared["state"], + controls, + SS.GOAL, + prepared["ped_xy"], + prepared["ped_vel"], + ).cpu().numpy() + for row, cost in zip(rows, costs): + row["expert_cost"] = float(cost) + return min( + rows, + key=lambda row: ( + row["expert_cost"], + int(row["candidate_id"]), + ), + ) + + +def _neutral_record( + neutral_id, + replica, + prepared, + chosen, + *, + step, + round_i=1, + selector, + repair_trigger, +): + result = chosen["result"] + if not _all_full_h_negative([chosen], expected=1): + raise ValueError("neutral execution requires one exact full-H negative") + controls = np.asarray(chosen["controls"], np.float32) + x0 = np.asarray(chosen["x0"], np.float32) + if tuple(controls.shape) != (10, 2) or not np.isfinite(controls).all(): + raise ValueError("neutral controls must be finite [10,2]") + if tuple(x0.shape) != (20,) or not np.isfinite(x0).all(): + raise ValueError("neutral x0 must be finite [20]") + return dict( + neutral_id=int(neutral_id), + population="D0", + semantic_label="neutral", + round=int(round_i), + scenario_id=int(replica.scenario_id), + gamma=float(replica.gamma), + step=int(step), + state=np.asarray(prepared["state"], np.float32), + hp10=np.asarray(prepared["hp10"].numpy(), np.float32), + low5=np.asarray(prepared["low"].numpy(), np.float32), + hist=np.asarray(prepared["hist"].numpy(), np.float32), + ped_xy=np.asarray(prepared["ped_xy"], np.float32), + ped_vel=np.asarray(prepared["ped_vel"], np.float32), + controls=controls, + x0=x0, + verifier_result=result, + verifier_y=int(result["y"]), + candidate_id=int(chosen["candidate_id"]), + parent_candidate_id=chosen.get("parent_candidate_id"), + query_id=chosen.get("query_id"), + sigma=float(chosen["sigma"]), + hp_margin=float(chosen["hp_margin"]), + expert_cost=( + None + if chosen.get("expert_cost") is None + else float(chosen["expert_cost"]) + ), + selector=str(selector), + repair_trigger=str(repair_trigger), + execution_source=( + f"kazuki_repair_neutral_{repair_trigger}_{selector}" + ), + train_eligible=False, + replay_default=False, + gp_eligible=False, + ) + + +def _save_neutral_records(path, records, *, round_i=1): + path = os.fspath(path) + for expected, record in enumerate(records): + if int(record["neutral_id"]) != expected: + raise AssertionError("neutral IDs are not dense") + if ( + record["semantic_label"] != "neutral" + or record["population"] != "D0" + or record["train_eligible"] + or record["replay_default"] + or record["gp_eligible"] + or int(record["verifier_y"]) != 0 + ): + raise AssertionError("invalid neutral execution record") + payload = dict( + version=1, + status="SFM_B1_NEUTRAL_ROUND_COMPLETE", + round=int(round_i), + records=records, + summary=dict(D0=len(records), train_eligible=0, gp_eligible=0), + ) + FA._save_torch(path, payload) + marker = dict( + status=payload["status"], + file=os.path.abspath(path), + sha256=FA._sha256_file(path), + **payload["summary"], + ) + FA._write_json(path + ".COMPLETE.json", marker) + return marker + + +def _gamma_balanced_gp( + phi_policy, + previous, + *, + gammas, + round_i, + ell, + cap, + lam, + phi_s, + device, + seed, +): + """Build a previous-round-only GP with an equal supported gamma quota.""" + gp = BR.RBFGP(float(ell), float(lam)) + empty = { + str(gamma): 0 for gamma in gammas + } + if previous is None: + return gp, [], dict( + requested_cap=int(cap), + effective_cap=0, + quota=0, + rotating_extra_gamma=None, + per_gamma=empty, + source_round=None, + population="previous executed D+ only", + ) + + groups = {} + for gamma_index, gamma in enumerate(gammas): + records = [ + (previous, row) + for row in previous.Dplus + if round( + float(previous.contexts[int(row["context_id"])]["gamma"]), + 8, + ) == round(float(gamma), 8) + ] + groups[float(gamma)] = BS.hierarchical_order( + records, int(seed) + gamma_index, + ) + quota = min( + int(cap) // len(gammas), + min(len(groups[float(gamma)]) for gamma in gammas), + ) + if quota == 0: + counts = { + str(gamma): len(groups[float(gamma)]) + for gamma in gammas + } + raise RuntimeError( + "previous-round D+ cannot support all declared gammas; " + f"per_gamma={counts}" + ) + selected = [ + record + for gamma in gammas + for record in groups[float(gamma)][:quota] + ] + rotation = (int(round_i) - 2) % len(gammas) + extra_gamma = None + if quota > 0 and len(selected) < int(cap): + for offset in range(len(gammas)): + gamma = float(gammas[(rotation + offset) % len(gammas)]) + if len(groups[gamma]) > quota: + selected.append(groups[gamma][quota]) + extra_gamma = gamma + break + + if selected: + feature_parts = [] + for start in range(0, len(selected), 256): + values = selected[start:start + 256] + hp10, low, hist, controls = BX._record_batch(values, device) + x0 = torch.as_tensor( + np.stack([row["x0"] for _, row in values]), + device=device, + ).float() + feature_parts.append(phi_policy.phi_s_from_x0( + controls, + phi_policy.ctx_from(hp10, low, hist), + x0, + s=float(phi_s), + )) + gp.set_buffer(torch.cat(feature_parts)) + identities = [ + (int(shard.round_i), int(row["window_id"])) + for shard, row in selected + ] + if len(identities) != len(set(identities)): + raise RuntimeError("dynamic gamma-balanced GP contains duplicates") + per_gamma = Counter( + str(previous.contexts[int(row["context_id"])]["gamma"]) + for _, row in selected + ) + return gp, identities, dict( + requested_cap=int(cap), + effective_cap=len(selected), + quota=int(quota), + rotating_extra_gamma=extra_gamma, + per_gamma={ + str(gamma): int(per_gamma[str(gamma)]) + for gamma in gammas + }, + source_round=int(previous.round_i), + population="previous executed D+ only; D0 excluded", + ) + + +@torch.no_grad() +def _calibrate_gp_beta( + phi_policy, gp, replicas, cfg, device, *, round_i, +): + live, batch = BX._stack_prepared(replicas, device) + windows, contexts, x0 = FA._keyed_windows( + phi_policy, + live, + batch, + K=cfg.K, + round_i=int(round_i), + step=-1, + source="beta_calibration", + seed=cfg.seed, + nfe=cfg.nfe, + temp=cfg.temp, + ) + features = FA._features_from_x0( + phi_policy, windows, contexts, x0, cfg.phi_s, + ) + score_vectors = [] + for replica, values in zip(live, features): + generator = torch.Generator(device=values.device).manual_seed( + FA._keyed_seed( + cfg.seed, + int(round_i), + replica.scenario_id, + f"{replica.gamma:.8f}", + "beta_order", + ) + ) + order = torch.randperm( + len(values), generator=generator, device=values.device, + ) + score_vectors.extend( + gp.sequential_score_vectors(values, order, cfg.B) + ) + beta, ess = BR.solve_beta( + score_vectors, target=cfg.ess_target, + ) + return float(beta), float(ess) + + +def _repair_trigger(replica, chosen): + current_trap = FA._trap(replica.states) + if current_trap: + return "trap_streak" + if chosen is None: + return "finite_B_NVP" + if KR.predicted_trap(replica.states, chosen["controls"][0]): + return "predicted_trap" + return None + + +def collect( + checkpoint, + *, + scenarios=DEFAULT_SCENARIOS, + gammas=DEFAULT_GAMMAS, + scene_profile="double_density_velocity_ood", + selector="margin", + device="cuda", + verifier_workers=16, + sample_seed=DEFAULT_SAMPLE_SEED, + audit_seed=DEFAULT_AUDIT_SEED, + ell=DEFAULT_ELL, + neutral_continuation=False, + round_i=1, + expected_checkpoint_sha256=EXPECTED_CHECKPOINT_SHA256, + previous_executed_path=None, + gp_cap=512, + ess_target=0.5, + verifier_executor=None, + T=180, + outdir, +): + """Collect immutable traces plus executed/query shards for one frozen model.""" + scenarios = tuple(map(int, scenarios)) + gammas = tuple(map(float, gammas)) + if not scenarios or len(set(scenarios)) != len(scenarios): + raise ValueError("scenarios must be distinct and nonempty") + if ( + not gammas + or len(set(gammas)) != len(gammas) + or any(value not in tuple(map(float, SS.GAMMAS)) for value in gammas) + ): + raise ValueError(f"gammas must be a distinct subset of {SS.GAMMAS}") + if selector not in ( + "margin", "progress_gated_margin", "safemppi_cost", + ): + raise ValueError( + "selector must be margin, progress_gated_margin, or " + "safemppi_cost" + ) + if scene_profile != "double_density_velocity_ood": + raise ValueError("repair audit is pinned to double-shift OOD") + if int(T) != 180: + raise ValueError("repair audit is scientifically pinned to T=180") + checkpoint = os.path.abspath(checkpoint) + outdir = os.path.abspath(outdir) + round_i = int(round_i) + if round_i < 1: + raise ValueError("round_i must be positive") + if int(gp_cap) < len(gammas): + raise ValueError("gp_cap must permit at least one row per gamma") + if os.path.exists(outdir): + raise FileExistsError(f"refusing to reuse output directory: {outdir}") + checkpoint_sha = FA._sha256_file(checkpoint) + if ( + expected_checkpoint_sha256 is None + or checkpoint_sha != str(expected_checkpoint_sha256) + ): + raise RuntimeError( + f"checkpoint SHA mismatch: expected " + f"{expected_checkpoint_sha256}, " + f"observed {checkpoint_sha}" + ) + previous = ( + None + if previous_executed_path is None + else OS.ExecutedRoundShard.load(previous_executed_path) + ) + if round_i == 1 and previous is not None: + raise ValueError("round 1 cannot have previous GP support") + if round_i > 1 and previous is None: + raise ValueError("round >1 requires previous executed D+ support") + if previous is not None and int(previous.round_i) != round_i - 1: + raise ValueError( + "previous executed shard must be from the immediately " + "preceding round" + ) + + environment = SS.scene_profile(scene_profile) + policy, _ = GPS.load_sfm_policy(checkpoint, device=device) + policy.eval() + phi_policy = copy.deepcopy(policy).eval() + for parameter in phi_policy.parameters(): + parameter.requires_grad_(False) + policy_hash = BX.policy_sha256(policy) + replicas = [ + BX.Replica( + scenario, + gamma, + n_ped=environment["n_ped"], + ped_speed_range=tuple(environment["ped_speed_range"]), + ) + for scenario in scenarios for gamma in gammas + ] + cfg = BX.ArmConfig( + name="diagnostic", + selector=selector, + alpha=0.0, + rounds=1, + K=16, + B=4, + T=int(T), + H=10, + W=2, + batch=128, + lr=1.0e-5, + ess_target=0.5, + nfe=8, + temp=1.0, + phi_s=0.9, + gp_lam=1.0e-2, + verifier_workers=int(verifier_workers), + smoke=False, + seed=int(audit_seed), + scene_profile=scene_profile, + ).validate() + if not 0.0 < float(ess_target) <= 1.0: + raise ValueError("ess_target must be in (0, 1]") + cfg = replace(cfg, ess_target=float(ess_target)) + gp, gp_ids, gp_selection = _gamma_balanced_gp( + phi_policy, + previous, + gammas=gammas, + round_i=round_i, + ell=float(ell), + cap=int(gp_cap), + lam=cfg.gp_lam, + phi_s=cfg.phi_s, + device=device, + seed=int(audit_seed) + round_i * 101, + ) + beta, calibrated_ess = _calibrate_gp_beta( + phi_policy, gp, replicas, cfg, device, round_i=round_i, + ) + + executed_shard = OS.ExecutedRoundShard(round_i) + query_shard = BS.RoundShard(round_i) + neutral_records = [] + traces = [] + counts = Counter() + sigma_pool, sigma_selected, ess_values = [], [], [] + trap_streaks = defaultdict(int) + + executor_scope = ( + ProcessPoolExecutor(max_workers=int(verifier_workers)) + if verifier_executor is None + else nullcontext(verifier_executor) + ) + with executor_scope as executor: + for step in range(int(T)): + live = [replica for replica in replicas if replica.alive] + live, batch = BX._stack_prepared(live, device) + if not live: + break + with torch.no_grad(): + windows, contexts, x0 = FA._keyed_windows( + policy, + live, + batch, + K=cfg.K, + round_i=round_i, + step=step, + source="K", + seed=int(sample_seed), + nfe=cfg.nfe, + temp=cfg.temp, + ) + features = FA._features_from_x0( + phi_policy, windows, contexts, x0, cfg.phi_s, + ) + windows_np = windows.detach().cpu().numpy() + x0_np = x0.detach().cpu().numpy() + + selected_by_context = [] + acquisition_by_context = [] + for context_index, replica in enumerate(live): + generator = torch.Generator( + device=features.device, + ).manual_seed(FA._keyed_seed( + int(audit_seed), + round_i, + replica.scenario_id, + f"{replica.gamma:.8f}", + step, + "acquisition", + )) + selected, acquisition = gp.sequential_acquire( + features[context_index], + cfg.B, + beta, + generator=generator, + ) + selected_by_context.append(list(map(int, selected))) + acquisition_by_context.append(acquisition) + sigma_pool.extend(map( + float, + gp.acquisition_sigma(features[context_index]) + .detach().cpu(), + )) + sigma_selected.extend( + float(row["chosen_sigma"]) for row in acquisition + ) + ess_values.extend(float(row["ess_norm"]) for row in acquisition) + + base_tasks = [] + for context_index, replica in enumerate(live): + prepared = replica.prepared + for candidate_id in selected_by_context[context_index]: + base_tasks.append(( + context_index, + candidate_id, + prepared["state"], + windows_np[context_index, candidate_id], + prepared["ped_xy"], + prepared["ped_vel"], + replica.gamma, + )) + base_results = list(executor.map(SM.verify_in_worker, base_tasks)) + counts["base_verifier_queries"] += len(base_tasks) + base_by_context = defaultdict(dict) + for context_index, candidate_id, result in base_results: + base_by_context[int(context_index)][int(candidate_id)] = result + + prepared_contexts = [] + repair_indices = [] + for context_index, replica in enumerate(live): + prepared = replica.prepared + prediction = SM.predict_pedestrians( + prepared["ped_xy"], prepared["ped_vel"], cfg.H, + ) + all_rows = [] + for candidate_id in range(cfg.K): + controls = windows_np[context_index, candidate_id] + segment = SM.rollout_positions(prepared["state"], controls) + all_rows.append(dict( + candidate_id=int(candidate_id), + controls=controls, + x0=x0_np[context_index, candidate_id], + segment=segment, + mode=BE.classify_candidate(segment, prediction), + )) + base_rows = [] + for acquisition_step, candidate_id in enumerate( + selected_by_context[context_index] + ): + source = all_rows[candidate_id] + base_rows.append(_query_row( + candidate_id, + source["controls"], + source["x0"], + base_by_context[context_index][candidate_id], + acquisition_step=acquisition_step, + sigma=acquisition_by_context[context_index][ + acquisition_step + ]["chosen_sigma"], + mode=source["mode"], + source="base_B", + )) + base_choice = _admissible( + base_rows, selector, prepared, replica.gamma, + ) + trigger = _repair_trigger(replica, base_choice) + prepared_contexts.append(dict( + all_rows=all_rows, + base_rows=base_rows, + base_choice=base_choice, + trigger=trigger, + repair_rows=[], + repair_diagnostics=None, + )) + if trigger is not None: + repair_indices.append(context_index) + + repair_tasks = [] + for context_index in repair_indices: + replica = live[context_index] + selected = selected_by_context[context_index] + guided, diagnostics = KR.same_latent_guided_controls( + policy, + contexts[context_index], + replica.prepared["state"], + replica.prepared["ped_xy"], + replica.prepared["ped_vel"], + replica.gamma, + x0[context_index, selected], + nfe=cfg.nfe, + collect_diagnostics=True, + ) + guided_np = guided.detach().cpu().numpy() + prepared_contexts[context_index][ + "repair_diagnostics" + ] = diagnostics + for acquisition_step, (candidate_id, controls) in enumerate( + zip(selected, guided_np) + ): + repair_id = cfg.K + int(candidate_id) + repair_tasks.append(( + context_index, + repair_id, + replica.prepared["state"], + controls, + replica.prepared["ped_xy"], + replica.prepared["ped_vel"], + replica.gamma, + )) + prepared_contexts[context_index]["repair_rows"].append( + dict( + repair_id=repair_id, + parent_candidate_id=int(candidate_id), + acquisition_step=int(acquisition_step), + controls=controls, + ) + ) + repair_results = list(executor.map(SM.verify_in_worker, repair_tasks)) + counts["repair_verifier_queries"] += len(repair_tasks) + repair_result_lookup = { + (int(context_index), int(candidate_id)): result + for context_index, candidate_id, result in repair_results + } + + for context_index, replica in enumerate(live): + prepared = replica.prepared + values = prepared_contexts[context_index] + sidecar_context_id = query_shard.add_context( + scenario_id=replica.scenario_id, + gamma=replica.gamma, + step=step, + state=prepared["state"], + hp10=prepared["hp10"].numpy(), + low5=prepared["low"].numpy(), + hist=prepared["hist"].numpy(), + ped_xy=prepared["ped_xy"], + ped_vel=prepared["ped_vel"], + ) + for row in values["base_rows"]: + _add_sidecar_query(query_shard, sidecar_context_id, row) + counts[f"base_{FA._result_label(row['result'])}"] += 1 + + repair_rows = [] + for repair in values["repair_rows"]: + parent_id = int(repair["parent_candidate_id"]) + source = values["all_rows"][parent_id] + repair_id = int(repair["repair_id"]) + row = _query_row( + repair_id, + repair["controls"], + source["x0"], + repair_result_lookup[(context_index, repair_id)], + acquisition_step=repair["acquisition_step"], + sigma=values["base_rows"][ + repair["acquisition_step"] + ]["sigma"], + mode=BE.classify_candidate( + SM.rollout_positions( + prepared["state"], repair["controls"], + ), + SM.predict_pedestrians( + prepared["ped_xy"], + prepared["ped_vel"], + cfg.H, + ), + ), + source="kazuki_guided_B", + parent_candidate_id=parent_id, + ) + repair_rows.append(row) + _add_sidecar_query(query_shard, sidecar_context_id, row) + counts[f"repair_{FA._result_label(row['result'])}"] += 1 + values["repair_rows"] = repair_rows + repaired_choice = ( + _admissible( + repair_rows, selector, prepared, replica.gamma, + ) + if values["trigger"] is not None else None + ) + neutral_choice = ( + _select_neutral( + repair_rows, selector, prepared, replica.gamma, + ) + if ( + bool(neutral_continuation) + and values["trigger"] is not None + and repaired_choice is None + and _all_full_h_negative(repair_rows) + ) + else None + ) + for row in repair_rows: + if row.get("query_id") is None: + continue + stored = query_shard.queries[int(row["query_id"])] + if "hp_margin" in row: + stored["hp_margin"] = float(row["hp_margin"]) + if "expert_cost" in row: + stored["expert_cost"] = float(row["expert_cost"]) + chosen = ( + repaired_choice or neutral_choice + if values["trigger"] is not None else values["base_choice"] + ) + + trace = dict( + round=round_i, + step=int(step), + scenario_id=int(replica.scenario_id), + gamma=float(replica.gamma), + state=prepared["state"].copy(), + next_state=prepared["state"].copy(), + ped_xy=prepared["ped_xy"].copy(), + ped_vel=prepared["ped_vel"].copy(), + all_K=values["all_rows"], + selected_ids=selected_by_context[context_index], + query_rows=values["base_rows"], + guided_query_rows=repair_rows, + acquisition=acquisition_by_context[context_index], + repair_trigger=values["trigger"], + repair_diagnostics=values["repair_diagnostics"], + repair_selected_id=None, + neutral_execution=False, + neutral_id=None, + executed_id=None, + executed_controls=None, + executed_x0=None, + executed_result=None, + execution_source=None, + trap_streak_before=int(trap_streaks[ + (replica.scenario_id, replica.gamma) + ]), + trap_event=False, + trap_fail_closed=False, + negative_reasons=[], + ) + if chosen is None: + replica.alive = False + replica.status = "repair_nvp" + trace["negative_reasons"].append("repair_no_admissible_B") + counts["repair_fail_closed"] += 1 + traces.append(trace) + continue + + is_repair = values["trigger"] is not None + is_neutral = ( + neutral_choice is not None and chosen is neutral_choice + ) + selected_x0 = np.asarray(chosen["x0"], np.float32) + controls = np.asarray(chosen["controls"], np.float32) + result = chosen["result"] + if ( + trap_streaks[(replica.scenario_id, replica.gamma)] + >= KR.TRAP_PATIENCE - 1 + and KR.predicted_trap(replica.states, controls[0]) + ): + replica.alive = False + replica.status = "trap_fail_closed" + trace.update( + trap_fail_closed=True, + negative_reasons=["predicted_third_consecutive_trap"], + ) + counts["trap_fail_closed"] += 1 + traces.append(trace) + continue + + execution_source = ( + ( + f"kazuki_repair_neutral_" + f"{values['trigger']}_{selector}" + ) + if is_neutral + else ( + f"kazuki_repair_{values['trigger']}_{selector}" + if is_repair else f"verified_{selector}" + ) + ) + window_id = None + neutral_record = None + if is_neutral: + neutral_record = _neutral_record( + len(neutral_records), + replica, + prepared, + chosen, + step=step, + round_i=round_i, + selector=selector, + repair_trigger=values["trigger"], + ) + neutral_records.append(neutral_record) + else: + executed_context_id = executed_shard.add_context( + scenario_id=replica.scenario_id, + gamma=replica.gamma, + step=step, + state=prepared["state"], + hp10=prepared["hp10"].numpy(), + low5=prepared["low"].numpy(), + hist=prepared["hist"].numpy(), + ped_xy=prepared["ped_xy"], + ped_vel=prepared["ped_vel"], + ) + window_id = executed_shard.add_executed_window( + executed_context_id, + controls, + selected_x0, + result, + execution_source=execution_source, + nvp_context=values["trigger"] == "finite_B_NVP", + candidate_id=chosen["candidate_id"], + acquisition_step=chosen["acquisition_step"], + sigma=chosen["sigma"], + hp_margin=chosen["hp_margin"], + mode=chosen["mode"], + ) + sidecar_query_id = chosen.get("query_id") + if sidecar_query_id is not None: + query_shard.mark_executed( + sidecar_query_id, + hp_margin=chosen["hp_margin"], + expert_cost=chosen.get("expert_cost"), + ) + query_stored = query_shard.queries[ + int(sidecar_query_id) + ] + query_stored["execution_role"] = ( + "neutral_continuation" + if is_neutral else "verified_execution" + ) + query_stored["neutral_id"] = ( + int(neutral_record["neutral_id"]) + if is_neutral else None + ) + + BX._advance(replica, controls[0]) + trap_event = FA._trap(replica.states) + trap_key = (replica.scenario_id, replica.gamma) + streak, stop = KR.next_trap_streak( + trap_streaks[trap_key], trap_event, + ) + trap_streaks[trap_key] = streak + collision, success, clearance = FA._post_action_terminal(replica) + if stop and replica.alive: + replica.alive = False + replica.status = "trap_fail_closed" + counts["trap_fail_closed"] += 1 + if is_neutral: + neutral_record.update( + next_state=replica.state.copy(), + trap_event=bool(trap_event), + trap_streak=int(streak), + collision_after_action=bool(collision), + success_after_action=bool(success), + clearance_after_action=float(clearance), + ) + counts["neutral_executions"] += 1 + else: + stored = executed_shard.windows[int(window_id)] + stored.update( + trap_event=bool(trap_event), + trap_streak=int(streak), + collision_after_action=bool(collision), + success_after_action=bool(success), + ) + counts["executed_training_windows"] += 1 + trace.update( + next_state=replica.state.copy(), + repair_selected_id=( + int(chosen["candidate_id"]) if is_repair else None + ), + neutral_execution=bool(is_neutral), + neutral_id=( + int(neutral_record["neutral_id"]) + if is_neutral else None + ), + executed_id=int(chosen["candidate_id"]), + executed_controls=controls, + executed_x0=selected_x0, + executed_result=result, + executed_label=( + "neutral" + if is_neutral else FA._result_label(result) + ), + executed_verifier_label=FA._result_label(result), + execution_source=execution_source, + window_id=( + None if window_id is None else int(window_id) + ), + trap_event=bool(trap_event), + trap_streak_after=int(streak), + trap_fail_closed=bool(stop), + collision_after_action=bool(collision), + success_after_action=bool(success), + clearance_after_action=float(clearance), + ) + counts["executed_windows"] += 1 + counts[f"source_{execution_source}"] += 1 + traces.append(trace) + + for replica in replicas: + if replica.alive: + replica.alive = False + replica.status = "timeout" + if BX.policy_sha256(policy) != policy_hash: + raise RuntimeError("policy changed during frozen repair collection") + executed_summary = executed_shard.validate() + query_summary = query_shard.validate() + if int(executed_summary["Dminus"]) != 0: + raise RuntimeError("repair collector executed a verifier-negative window") + if len(neutral_records) != int(counts["neutral_executions"]): + raise RuntimeError("neutral execution accounting mismatch") + if int(counts["executed_training_windows"]) != int( + executed_summary["D"] + ): + raise RuntimeError("executed training-window accounting mismatch") + if int(counts["executed_windows"]) != ( + int(executed_summary["D"]) + len(neutral_records) + ): + raise RuntimeError("total executed-action accounting mismatch") + if not bool(neutral_continuation) and neutral_records: + raise RuntimeError("default fail-closed arm produced neutral records") + if int(counts["base_verifier_queries"]) != len(query_shard.contexts) * cfg.B: + raise RuntimeError("base B=4 accounting mismatch") + + os.makedirs(outdir) + executed_path = os.path.join(outdir, "executed_round.pt") + query_path = os.path.join(outdir, "query_sidecar.pt") + neutral_path = os.path.join(outdir, "neutral_round.pt") + executed_manifest = executed_shard.save(executed_path) + query_manifest = query_shard.save(query_path) + neutral_manifest = _save_neutral_records( + neutral_path, neutral_records, round_i=round_i, + ) + outcomes = [dict( + scenario_id=int(replica.scenario_id), + gamma=float(replica.gamma), + status=str(replica.status), + success=replica.status == "success", + collision=replica.status == "collision", + timeout=replica.status == "timeout", + repair_nvp=replica.status == "repair_nvp", + trap_fail_closed=replica.status == "trap_fail_closed", + steps=len(replica.controls), + minimum_clearance=float(replica.minimum_clearance), + ) for replica in replicas] + bundle = dict( + version=1, + status=STATUS, + round=round_i, + source=FA._source(), + checkpoint=checkpoint, + checkpoint_sha256=checkpoint_sha, + scene_profile=scene_profile, + environment=environment, + scenarios=list(scenarios), + gammas=list(gammas), + selector=selector, + neutral_continuation=bool(neutral_continuation), + protocol=dict( + K=cfg.K, + B=cfg.B, + H=cfg.H, + T=int(T), + ell=float(ell), + gp_cap=int(gp_cap), + gp_buffer_ids=gp_ids, + gp_selection=gp_selection, + beta=float(beta), + calibrated_ess_over_K=float(calibrated_ess), + realized_ess_over_K=float(np.mean(ess_values)), + gp_diagnostics=gp.diagnostics(), + acquisition=BR.acquisition_diagnostics( + sigma_pool, sigma_selected, + ), + repair=KR.manifest(), + D_exec=( + "one exact-positive executed H10 window per executed context" + ), + D_query=( + "all resolved base and repair B labels; audit-only by default" + ), + D_neutral=( + "guided full-H verifier-negative actions actually executed " + "after an NVP/trap repair trigger; isolated from " + "training D+/D-/GP/replay and retained as exact-negative " + "in the audit query sidecar" + ), + GP_population="unchanged: prior executed D+ only", + ), + sample_seed=int(sample_seed), + audit_seed=int(audit_seed), + counts=dict(counts), + outcomes=outcomes, + executed_shard=executed_manifest, + query_sidecar=query_manifest, + neutral_shard=neutral_manifest, + traces=traces, + ) + trace_path = os.path.join(outdir, "repair_trace.pt") + FA._save_torch(trace_path, bundle) + marker = dict( + status=STATUS, + source=bundle["source"], + trace_path=os.path.abspath(trace_path), + trace_sha256=FA._sha256_file(trace_path), + checkpoint_sha256=checkpoint_sha, + round=round_i, + selector=selector, + neutral_continuation=bool(neutral_continuation), + counts=dict(counts), + outcomes=outcomes, + executed_shard=executed_manifest, + query_sidecar=query_manifest, + neutral_shard=neutral_manifest, + gp_buffer_ids=gp_ids, + gp_selection=gp_selection, + gp_diagnostics=gp.diagnostics(), + acquisition=bundle["protocol"]["acquisition"], + ) + FA._write_json(os.path.join(outdir, "COMPLETE.json"), marker) + return trace_path + + +def main(argv=None): + parser = argparse.ArgumentParser() + parser.add_argument("--checkpoint", required=True) + parser.add_argument("--outdir", required=True) + parser.add_argument( + "--scenarios", type=int, nargs="+", default=DEFAULT_SCENARIOS, + ) + parser.add_argument( + "--gammas", type=float, nargs="+", default=DEFAULT_GAMMAS, + ) + parser.add_argument( + "--scene-profile", default="double_density_velocity_ood", + ) + parser.add_argument( + "--selector", choices=( + "margin", "progress_gated_margin", "safemppi_cost", + ), default="margin", + ) + parser.add_argument("--device", default="cuda") + parser.add_argument("--verifier-workers", type=int, default=16) + parser.add_argument("--sample-seed", type=int, default=DEFAULT_SAMPLE_SEED) + parser.add_argument("--audit-seed", type=int, default=DEFAULT_AUDIT_SEED) + parser.add_argument("--ell", type=float, default=DEFAULT_ELL) + parser.add_argument("--ess-target", type=float, default=0.5) + parser.add_argument("--neutral-continuation", action="store_true") + args = parser.parse_args(argv) + collect( + args.checkpoint, + scenarios=args.scenarios, + gammas=args.gammas, + scene_profile=args.scene_profile, + selector=args.selector, + device=args.device, + verifier_workers=args.verifier_workers, + sample_seed=args.sample_seed, + audit_seed=args.audit_seed, + ell=args.ell, + ess_target=args.ess_target, + neutral_continuation=args.neutral_continuation, + T=180, + outdir=args.outdir, + ) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/sfm_b1_kazuki_repair_sixrow.py b/overnight_run_07_12_sfm/sfm_b1_kazuki_repair_sixrow.py new file mode 100644 index 0000000..32570ff --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_kazuki_repair_sixrow.py @@ -0,0 +1,720 @@ +"""Two-episode, six-row mechanism comparison for Kazuki-repaired gathering. + +Rows are episode-major: + +1. raw pretrained deployment with exact verifier geometry; +2. locked Kazuki generate--guide--refine with separate guidance arrows; +3. same-latent B4 repair gathering; + +then the same three rows for the second fixed OOD episode. Rendering consumes +an immutable trace bundle and never reruns a controller or verifier. +""" +from __future__ import annotations + +import argparse +from concurrent.futures import ProcessPoolExecutor +import json +import os + +import matplotlib +matplotlib.use("Agg") +import matplotlib.animation as animation +import matplotlib.pyplot as plt +from matplotlib.lines import Line2D +import numpy as np +import torch + +import _paths # noqa: F401 +import grid_policy_sfm as GPS +import sfm_b1_density_viz as DV +import sfm_b1_eval as BE +import sfm_b1_full_episode_viz as FV +import sfm_b1_kazuki_repair_audit as RA +import sfm_b1_viz as BV +import sfm_kazuki as KZ +import sfm_metrics2 as SM +import sfm_scene as SS + + +STATUS = "SFM_B1_KAZUKI_REPAIR_SIXROW_TRACE_COMPLETE" +RENDER_STATUS = "SFM_B1_KAZUKI_REPAIR_SIXROW_RENDER_COMPLETE" +TRUE_BLUE = "#0057FF" +TRUE_RED = "#D62728" +MAGENTA = "#CC79A7" +CYAN = "#00A6D6" +GREEN = "#009E73" +DEFAULT_EPISODES = (250_001, 250_003) +DEFAULT_GAMMAS = (0.1, 0.5, 1.0) + + +def _write_json(path, payload): + path = os.path.abspath(path) + os.makedirs(os.path.dirname(path), exist_ok=True) + temporary = path + ".tmp" + with open(temporary, "w") as stream: + json.dump(payload, stream, indent=2, sort_keys=True, allow_nan=False) + os.replace(temporary, path) + + +def _save_torch(path, payload): + path = os.path.abspath(path) + os.makedirs(os.path.dirname(path), exist_ok=True) + temporary = path + ".tmp" + torch.save(payload, temporary) + os.replace(temporary, path) + + +def _collect_raw(policy, episodes, gammas, environment, *, device, T, sample_seed): + runs = {} + pending = [] + for episode in episodes: + for gamma in gammas: + run = BE.raw_rollout( + policy, + episode, + gamma, + device=device, + T=int(T), + n_ped=int(environment["n_ped"]), + temp=1.0, + nfe=8, + ped_speed_range=tuple(environment["ped_speed_range"]), + sample_seed=int(sample_seed), + collect_trace=True, + ) + runs[(int(episode), float(gamma))] = run + for trace_index, trace in enumerate(run["trace"]): + pending.append(( + len(pending), + 0, + trace["state"], + trace["controls"], + trace["ped_xy"], + trace["ped_vel"], + float(gamma), + int(episode), + int(trace_index), + )) + + worker_tasks = [row[:7] for row in pending] + workers = min(32, max(1, os.cpu_count() or 1)) + with ProcessPoolExecutor(max_workers=workers) as executor: + results = list(executor.map(SM.verify_in_worker, worker_tasks)) + for pending_row, result_row in zip(pending, results): + _, _, result = result_row + episode = pending_row[-2] + trace_index = pending_row[-1] + gamma = float(pending_row[6]) + runs[(episode, gamma)]["trace"][trace_index][ + "verifier_result" + ] = result + return runs + + +def _collect_kazuki( + policy, episodes, gammas, environment, *, device, T, sample_seed, +): + config = KZ.KazukiConfig( + safe_coefs=(0.3,), + goal_coef=0.5, + safe_coef_gamma_span=0.0, + goal_coef_gamma_span=0.0, + ).validate() + return { + (int(episode), float(gamma)): KZ.kazuki_sfm_deploy( + policy, + episode=int(episode), + gamma=float(gamma), + cfg=config, + n_ped=int(environment["n_ped"]), + T=int(T), + device=device, + ped_speed_range=tuple(environment["ped_speed_range"]), + sample_seed=int(sample_seed), + collect_diagnostics=True, + ) + for episode in episodes for gamma in gammas + } + + +def collect( + checkpoint, + output_dir, + *, + episodes=DEFAULT_EPISODES, + gammas=DEFAULT_GAMMAS, + scene_profile="double_density_velocity_ood", + selector="margin", + device="cuda", + verifier_workers=16, + sample_seed=700_000, + audit_seed=20260730, + ell=RA.DEFAULT_ELL, + neutral_continuation=False, + T=180, +): + episodes = tuple(map(int, episodes)) + gammas = tuple(map(float, gammas)) + if len(episodes) != 2 or len(set(episodes)) != 2: + raise ValueError("six-row comparison requires two distinct episodes") + if gammas != DEFAULT_GAMMAS: + raise ValueError(f"six-row comparison requires gammas={DEFAULT_GAMMAS}") + if int(T) != 180: + raise ValueError("six-row comparison is scientifically pinned to T=180") + output_dir = os.path.abspath(output_dir) + if os.path.exists(output_dir): + raise FileExistsError(f"refusing to reuse output directory: {output_dir}") + os.makedirs(output_dir) + environment = SS.scene_profile(scene_profile) + policy, _ = GPS.load_sfm_policy(checkpoint, device=device) + policy.eval() + + raw = _collect_raw( + policy, + episodes, + gammas, + environment, + device=device, + T=T, + sample_seed=sample_seed, + ) + kazuki = _collect_kazuki( + policy, + episodes, + gammas, + environment, + device=device, + T=T, + sample_seed=sample_seed, + ) + repair_dir = os.path.join(output_dir, "repair") + repair_path = RA.collect( + checkpoint, + scenarios=episodes, + gammas=gammas, + scene_profile=scene_profile, + selector=selector, + device=device, + verifier_workers=verifier_workers, + sample_seed=sample_seed, + audit_seed=audit_seed, + ell=ell, + neutral_continuation=neutral_continuation, + T=T, + outdir=repair_dir, + ) + repair = torch.load( + repair_path, map_location="cpu", weights_only=False, + ) + bundle = dict( + version=1, + status=STATUS, + source=RA.FA._source(), + checkpoint=os.path.abspath(checkpoint), + checkpoint_sha256=RA.FA._sha256_file(checkpoint), + episodes=list(episodes), + gammas=list(gammas), + scene_profile=scene_profile, + environment=environment, + selector=selector, + neutral_continuation=bool(neutral_continuation), + sample_seed=int(sample_seed), + audit_seed=int(audit_seed), + raw=raw, + kazuki=kazuki, + repair=repair, + method_semantics=dict( + raw=( + "pretrained temp=1,NFE=8 raw policy; exact verifier geometry " + "is audit-only and never alters execution" + ), + kazuki=( + "locked pretrained prior, safe_coef=.3, goal_coef=.5, full " + "generate-guide-refine comparator" + ), + repair=( + "K=16, RBF B=4; same-latent locked guidance only on NVP/trap; " + "guided B is reverified; optional exact-negative neutral " + "continuation; no raw fallback or privileged MPC" + ), + ), + ) + trace_path = os.path.join(output_dir, "sixrow_trace.pt") + _save_torch(trace_path, bundle) + _write_json(os.path.join(output_dir, "TRACE_COMPLETE.json"), dict( + status=STATUS, + source=bundle["source"], + trace_path=os.path.abspath(trace_path), + trace_sha256=RA.FA._sha256_file(trace_path), + checkpoint_sha256=bundle["checkpoint_sha256"], + episodes=list(episodes), + gammas=list(gammas), + selector=selector, + neutral_continuation=bool(neutral_continuation), + repair_complete=os.path.join(repair_dir, "COMPLETE.json"), + )) + return trace_path + + +def _repair_index(traces): + output = {} + for trace in traces: + key = (int(trace["scenario_id"]), round(float(trace["gamma"]), 8)) + output.setdefault(key, {})[int(trace["step"])] = trace + return output + + +def _clamped_trace(run, step): + traces = list(run.get("trace") or ()) + if not traces: + raise ValueError("run has no mechanism trace") + return traces[min(max(int(step), 0), len(traces) - 1)] + + +def _draw_raw(axis, run, gamma, step): + trace = _clamped_trace(run, step) + DV.draw_method_panel( + axis, + "selected", + run, + gamma, + int(step), + verifier_result=trace["verifier_result"], + ) + + +def _draw_kazuki(axis, run, gamma, step): + DV.draw_method_panel( + axis, + "kazuki", + run, + gamma, + int(step), + guidance_scale=3.0, + guidance_cap=1.8, + ) + + +def _branch_path(trace): + result = trace.get("executed_result") + if not result: + return None + path = np.asarray(result.get("segment", ()), float) + return path if path.shape == (11, 2) else None + + +def _draw_repair_history(axis, rows, step): + states = [] + for value in sorted(key for key in rows if key <= int(step)): + trace = rows[value] + if trace.get("executed_result") is None: + continue + path = _branch_path(trace) + if path is not None: + branch_color = ( + MAGENTA + if str(trace.get("execution_source", "")).startswith( + "kazuki_repair_" + ) + else TRUE_BLUE + ) + axis.plot( + path[:, 0], path[:, 1], + color=branch_color, + lw=.7, + marker=".", + ms=1.1, + alpha=.32, + zorder=3, + ) + states.append(np.asarray(trace["state"], float)[:2]) + states.append(np.asarray(trace["next_state"], float)[:2]) + if states: + path = np.asarray(states)[::2] + final = np.asarray(states[-1])[None] + path = np.concatenate([path, final], axis=0) + axis.plot( + path[:, 0], path[:, 1], + color="#111111", + lw=1.25, + marker=".", + ms=1.4, + zorder=10, + ) + + +def _draw_query_path(axis, trace, row, *, color, linewidth, alpha): + path = np.asarray( + row["result"].get( + "segment", + SM.rollout_positions(trace["state"], row["controls"]), + ), + float, + ) + axis.plot( + path[:, 0], + path[:, 1], + color=color, + lw=float(linewidth), + marker=".", + ms=1.35, + alpha=float(alpha), + zorder=6, + ) + result = row["result"] + if not (result.get("resolved") and int(result.get("y", 0)) == 1): + axis.plot( + path[-1, 0], path[-1, 1], + "x", color=TRUE_RED, ms=4.4, mew=1.0, zorder=8, + ) + + +def _draw_repair(axis, rows, step): + available = [value for value in rows if value <= int(step)] + current_step = max(available) if available else min(rows) + trace = rows[current_step] + BV._draw_common(axis, trace, nominal_levels=False) + _draw_repair_history(axis, rows, current_step) + for row in trace["all_K"]: + path = np.asarray(row["segment"], float) + axis.plot( + path[:, 0], path[:, 1], + color=BV.GRAY, lw=.35, marker=".", ms=.9, + alpha=.18, zorder=2, + ) + for row in trace["query_rows"]: + _draw_query_path( + axis, trace, row, color=GREEN, linewidth=.95, alpha=.85, + ) + for row in trace.get("guided_query_rows", ()): + selected = ( + trace.get("repair_selected_id") is not None + and int(row["candidate_id"]) == int(trace["repair_selected_id"]) + ) + _draw_query_path( + axis, + trace, + row, + color=MAGENTA, + linewidth=2.0 if selected else 1.15, + alpha=.98 if selected else .76, + ) + if ( + trace.get("executed_result") is not None + and int(trace["executed_result"].get("y", 0)) == 1 + ): + audit = DV.checked_verifier_levels( + trace, dict(result=trace["executed_result"]), H=10, + ) + DV._draw_verifier_geometry(axis, audit) + elif trace.get("executed_result") is None: + position = np.asarray(trace["state"], float)[:2] + axis.plot( + position[0], position[1], + marker="o", ms=9, mfc="none", mec=TRUE_RED, + mew=1.5, zorder=12, + ) + axis.plot( + SS.GOAL[0], SS.GOAL[1], "*", + color="#F0E442", mec="#333333", ms=8, zorder=15, + ) + DV._set_clean_axis(axis) + + +def _layout(bundle): + gammas = tuple(map(float, bundle["gammas"])) + figure = plt.figure(figsize=(3.05 * len(gammas) + 4.2, 16.6)) + grid = figure.add_gridspec( + 6, + len(gammas) + 1, + width_ratios=[1.0] * len(gammas) + [1.22], + left=.025, + right=.985, + bottom=.025, + top=.985, + wspace=.025, + hspace=.035, + ) + axes = np.empty((6, len(gammas)), dtype=object) + sides = [] + for row in range(6): + for column in range(len(gammas)): + axes[row, column] = figure.add_subplot(grid[row, column]) + side = figure.add_subplot(grid[row, -1]) + side.set_axis_off() + sides.append(side) + return figure, axes, sides, gammas + + +def _side_text( + bundle, + row, + frame_index, + frame_count, + simulator_step, + repair_index, +): + episode = int(bundle["episodes"][row // 3]) + method = row % 3 + labels = ( + ( + "Raw pretrained r0", + "temperature 1 · NFE 8", + "verifier geometry is audit-only", + ), + ( + "Locked Kazuki comparator", + "safe=.3 · goal=.5", + "cyan=goal · magenta=safety", + ), + ( + ( + f"Active B4 neutral repair · {bundle['selector']}" + if bundle.get("neutral_continuation") + else f"Active B4 repair · {bundle['selector']}" + ), + "base B=green · guided B=magenta", + "no independent raw / privileged MPC", + ), + ) + lines = [ + f"episode {episode}", + *labels[method], + ] + if method == 2: + outcomes = { + ( + int(value["scenario_id"]), + round(float(value["gamma"]), 8), + ): value + for value in bundle["repair"]["outcomes"] + } + local_status = [] + for gamma in bundle["gammas"]: + key = (episode, round(float(gamma), 8)) + traces = repair_index[key] + final_step = max(traces) + local_step = min(int(simulator_step), final_step) + terminal = ( + str(outcomes[key]["status"]).replace("_", " ") + if int(simulator_step) >= final_step else "" + ) + suffix = f" · {terminal}" if terminal else "" + local_status.append( + f"γ={float(gamma):g}: t={local_step}{suffix}" + ) + lines.extend(("", *local_status)) + if row == 0: + columns = ", ".join(f"{value:g}" for value in bundle["gammas"]) + lines.extend(( + "", + f"columns γ: {columns}", + f"frame {frame_index}/{frame_count - 1}", + f"simulator step {simulator_step}", + )) + return "\n".join(lines) + + +def _maximum_step(bundle, repair_index): + maximum = 0 + for group in ("raw", "kazuki"): + for run in bundle[group].values(): + maximum = max(maximum, len(run.get("trace") or ()) - 1) + for rows in repair_index.values(): + maximum = max(maximum, max(rows)) + return maximum + + +def draw_frame(bundle, frame_index, frames, *, layout=None): + if bundle.get("status") != STATUS: + raise ValueError("not a completed repair six-row bundle") + if layout is None: + layout = _layout(bundle) + figure, axes, sides, gammas = layout + step = int(frames[int(frame_index)]) + repair_index = _repair_index(bundle["repair"]["traces"]) + for row in range(6): + sides[row].clear() + sides[row].set_axis_off() + sides[row].text( + .02, + .96, + _side_text( + bundle, + row, + frame_index, + len(frames), + step, + repair_index, + ), + ha="left", + va="top", + fontsize=8.0, + linespacing=1.42, + ) + for episode_index, episode in enumerate(bundle["episodes"]): + row0 = 3 * episode_index + for column, gamma in enumerate(gammas): + for row in range(row0, row0 + 3): + axes[row, column].clear() + key = (int(episode), float(gamma)) + _draw_raw(axes[row0, column], bundle["raw"][key], gamma, step) + _draw_kazuki( + axes[row0 + 1, column], + bundle["kazuki"][key], + gamma, + step, + ) + _draw_repair( + axes[row0 + 2, column], + repair_index[(int(episode), round(float(gamma), 8))], + step, + ) + return layout + + +def _legend(neutral_continuation=False): + return [ + Line2D([], [], color=TRUE_BLUE, lw=1.4, label=r"base executed exact-positive $D^+$"), + Line2D([], [], color="#111111", lw=1.2, label="executed first-action path"), + Line2D([], [], color=GREEN, lw=1.1, label="base RBF B=4 query"), + Line2D( + [], [], color=MAGENTA, lw=1.5, + label=( + "guided B repair / neutral executed branch" + if neutral_continuation + else "guided B repair / repaired executed branch" + ), + ), + Line2D([], [], color=TRUE_RED, marker="x", lw=0, label="exact rejected endpoint"), + Line2D([], [], color=GREEN, lw=.7, label="executed verifier levels h=1..10"), + Line2D([], [], color=CYAN, lw=2.2, label=r"Kazuki integrated $\nabla$ goal"), + Line2D([], [], color=MAGENTA, lw=2.2, label=r"Kazuki integrated $\nabla$ safety"), + ] + + +def render(trace_path, output_dir, *, fps=5, frame_stride=2, dpi=105): + output_dir = os.path.abspath(output_dir) + os.makedirs(output_dir, exist_ok=True) + bundle = torch.load( + trace_path, map_location="cpu", weights_only=False, + ) + repair_index = _repair_index(bundle["repair"]["traces"]) + maximum = _maximum_step(bundle, repair_index) + frames = list(range(0, maximum + 1, int(frame_stride))) + if frames[-1] != maximum: + frames.append(maximum) + layout = _layout(bundle) + figure = layout[0] + figure.legend( + handles=_legend(bool(bundle.get("neutral_continuation"))), + loc="lower right", + bbox_to_anchor=(.985, .012), + frameon=False, + fontsize=7.0, + ) + draw_frame(bundle, len(frames) - 1, frames, layout=layout) + last_frame = os.path.join(output_dir, "sixrow_last_frame.png") + figure.savefig(last_frame, dpi=170) + + def update(frame_index): + draw_frame(bundle, frame_index, frames, layout=layout) + return [] + + movie = animation.FuncAnimation( + figure, + update, + frames=range(len(frames)), + interval=1000 / int(fps), + blit=False, + ) + mp4 = os.path.join(output_dir, "sixrow_comparison.mp4") + movie.save( + mp4, + writer=animation.FFMpegWriter(fps=int(fps), bitrate=5200), + dpi=int(dpi), + ) + plt.close(figure) + report = dict( + status=RENDER_STATUS, + trace_source=bundle.get("source"), + renderer_source=RA.FA._source(), + trace_path=os.path.abspath(trace_path), + trace_sha256=RA.FA._sha256_file(trace_path), + frames=frames, + frame_count=len(frames), + fps=int(fps), + frame_stride=int(frame_stride), + mp4=os.path.abspath(mp4), + last_frame_png=os.path.abspath(last_frame), + episodes=list(map(int, bundle["episodes"])), + gammas=list(map(float, bundle["gammas"])), + selector=bundle["selector"], + neutral_continuation=bool(bundle.get("neutral_continuation")), + ) + _write_json(os.path.join(output_dir, "RENDER_COMPLETE.json"), report) + return report + + +def main(argv=None): + parser = argparse.ArgumentParser() + subparsers = parser.add_subparsers(dest="command", required=True) + collect_parser = subparsers.add_parser("collect") + collect_parser.add_argument("--checkpoint", required=True) + collect_parser.add_argument("--output-dir", required=True) + collect_parser.add_argument( + "--episodes", type=int, nargs=2, default=DEFAULT_EPISODES, + ) + collect_parser.add_argument( + "--gammas", type=float, nargs=3, default=DEFAULT_GAMMAS, + ) + collect_parser.add_argument( + "--scene-profile", default="double_density_velocity_ood", + ) + collect_parser.add_argument( + "--selector", choices=("margin", "safemppi_cost"), default="margin", + ) + collect_parser.add_argument("--device", default="cuda") + collect_parser.add_argument("--verifier-workers", type=int, default=16) + collect_parser.add_argument("--sample-seed", type=int, default=700_000) + collect_parser.add_argument("--audit-seed", type=int, default=20260730) + collect_parser.add_argument("--ell", type=float, default=RA.DEFAULT_ELL) + collect_parser.add_argument( + "--neutral-continuation", action="store_true", + ) + + render_parser = subparsers.add_parser("render") + render_parser.add_argument("--trace", required=True) + render_parser.add_argument("--output-dir", required=True) + render_parser.add_argument("--fps", type=int, default=5) + render_parser.add_argument("--frame-stride", type=int, default=2) + render_parser.add_argument("--dpi", type=int, default=105) + args = parser.parse_args(argv) + if args.command == "collect": + collect( + args.checkpoint, + args.output_dir, + episodes=args.episodes, + gammas=args.gammas, + scene_profile=args.scene_profile, + selector=args.selector, + device=args.device, + verifier_workers=args.verifier_workers, + sample_seed=args.sample_seed, + audit_seed=args.audit_seed, + ell=args.ell, + neutral_continuation=args.neutral_continuation, + T=180, + ) + else: + render( + args.trace, + args.output_dir, + fps=args.fps, + frame_stride=args.frame_stride, + dpi=args.dpi, + ) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/sfm_b1_neutral_multiround.py b/overnight_run_07_12_sfm/sfm_b1_neutral_multiround.py new file mode 100644 index 0000000..2f19526 --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_neutral_multiround.py @@ -0,0 +1,2187 @@ +"""Small-lineage multi-round SFM expansion with paired trigger probes. + +This additive study keeps the existing max-margin collector semantics while +changing only the experiment schedule: + +* two scenario seeds x all seven gammas per macro-round; +* previous-round executed D+ only in the RBF GP; +* whole-D+ followed by isolated whole-D0 CFM replay; +* one occurrence of every row per inner pass and one Adam step per pass; +* fixed-context, fixed-latent exact-H10 probes before and after each phase. + +D0 retains its exact verifier label y=0. It is a separate teacher population +and never enters D+, the RBF GP, or the ordinary replay population. +""" +from __future__ import annotations + +import argparse +from collections import defaultdict +from concurrent.futures import ProcessPoolExecutor +from dataclasses import asdict, dataclass +import hashlib +import json +import math +import multiprocessing as mp +import os +import time + +import numpy as np +import torch + +import _paths # noqa: F401 +import grid_policy_sfm as GPS +import sfm_b1_cost as BC +import sfm_b1_eval as BE +import sfm_b1_expand as BX +import sfm_b1_full_episode_audit as FA +import sfm_b1_kazuki_repair_audit as RA +import sfm_b1_neutral_teacher_sanity as NS +import sfm_b1_offline_eval as OE +import sfm_b1_offline_store as OS +import sfm_b1_r2_alpha_replay as R2 +import sfm_b1_rbf as BR +import sfm_b1_store as BS +import sfm_metrics2 as SM +import sfm_protocol as SP +import sfm_scene as SS + + +STATUS = "SFM_B1_NEUTRAL_MULTIROUND_COMPLETE" +ROUND_STATUS = "SFM_B1_NEUTRAL_MULTIROUND_ROUND_COMPLETE" +PROBE_STATUS = "SFM_B1_PAIRED_TRIGGER_PROBE_COMPLETE" +DEFAULT_SCENARIO_EP0 = 260_000 +DEFAULT_EVAL_EP0 = 270_000 +DEFAULT_PROBE_PER_GAMMA = 4 +DEFAULT_GP_CAP = 512 +DEFAULT_LR = 3.0e-5 +DEFAULT_INNER_STEPS = 4 + + +@dataclass(frozen=True) +class StudyConfig: + name: str + rounds: int = 2 + scenarios_per_round: int = 2 + lr: float = DEFAULT_LR + inner_steps: int = DEFAULT_INNER_STEPS + batch: int = 128 + K: int = 16 + B: int = 4 + H: int = 10 + T: int = 180 + nfe: int = 8 + temp: float = 1.0 + phi_s: float = 0.9 + ell: float = RA.DEFAULT_ELL + gp_cap: int = DEFAULT_GP_CAP + gp_lambda: float = 1.0e-2 + ess_target: float = 0.5 + selector: str = "margin" + encoder_lr_ratio: float = 0.0 + neutral_replay: bool = True + alpha: float = 0.0 + scene_profile: str = "double_density_velocity_ood" + sample_seed: int = 700_000 + audit_seed: int = 2_026_073_0 + train_seed: int = 2_026_073_1 + probe_seed: int = 2_026_073_2 + + def validate(self): + if int(self.rounds) < 1: + raise ValueError("rounds must be positive") + if int(self.scenarios_per_round) != 2: + raise ValueError("this study requires exactly two scenarios/round") + if not math.isfinite(float(self.lr)) or float(self.lr) <= 0.0: + raise ValueError("lr must be finite and positive") + if int(self.inner_steps) < 1: + raise ValueError("inner_steps must be positive") + if not 0.0 < float(self.ess_target) <= 1.0: + raise ValueError("ess_target must be in (0, 1]") + if int(self.batch) != 128: + raise ValueError("canonical microbatch size is 128") + if ( + int(self.K) != 16 + or int(self.B) != 4 + or int(self.H) != 10 + or int(self.T) != 180 + ): + raise ValueError("K/B/H/T protocol changed") + if self.selector not in ("margin", "progress_gated_margin"): + raise ValueError("unsupported neutral-study selector") + if float(self.encoder_lr_ratio) not in (0.0, 0.1): + raise ValueError("encoder_lr_ratio must be 0 (frozen) or 0.1") + if not isinstance(self.neutral_replay, bool): + raise ValueError("neutral_replay must be boolean") + if float(self.alpha) != 0.0: + raise ValueError("this study is pinned to alpha=0") + if self.scene_profile != "double_density_velocity_ood": + raise ValueError("scene profile changed") + return self + + +class _NeutralHolder: + def __init__(self, round_i): + self.round_i = int(round_i) + self.contexts = [] + self.windows = [] + + +def _write_json(path, payload): + path = os.path.abspath(os.fspath(path)) + os.makedirs(os.path.dirname(path), exist_ok=True) + temporary = path + ".tmp" + with open(temporary, "w") as stream: + json.dump(payload, stream, indent=2, sort_keys=True, allow_nan=False) + os.replace(temporary, path) + + +def _sha256_jsonable(value): + def convert(item): + if isinstance(item, np.ndarray): + return { + "dtype": str(item.dtype), + "shape": list(item.shape), + "bytes": hashlib.sha256(item.tobytes()).hexdigest(), + } + if isinstance(item, np.generic): + return item.item() + if isinstance(item, dict): + return { + str(key): convert(item[key]) + for key in sorted(item, key=str) + } + if isinstance(item, (list, tuple)): + return [convert(row) for row in item] + return item + + encoded = json.dumps( + convert(value), sort_keys=True, separators=(",", ":"), + ).encode() + return hashlib.sha256(encoded).hexdigest() + + +def _compact_mass(accounting): + return { + "total": float(accounting["total"]), + "gamma": { + str(key): float(value) + for key, value in accounting["gamma"].items() + }, + "cells": len(accounting["cells"]), + "contexts": len(accounting["contexts"]), + } + + +def _neutral_records(path): + payload = torch.load(path, map_location="cpu", weights_only=False) + if payload.get("status") != "SFM_B1_NEUTRAL_ROUND_COMPLETE": + raise ValueError("not an authenticated neutral-round payload") + holder = _NeutralHolder(payload["round"]) + records = [] + for expected, source in enumerate(payload["records"]): + result = source.get("verifier_result", {}) + if ( + int(source["neutral_id"]) != expected + or source["population"] != "D0" + or source["semantic_label"] != "neutral" + or int(source["verifier_y"]) != 0 + or not result.get("resolved") + or int(result.get("y", -1)) != 0 + or not result.get("full_h") + or int(result.get("terminal_step", -1)) != SP.H + or source["train_eligible"] + or source["replay_default"] + or source["gp_eligible"] + ): + raise RuntimeError("D0 semantics changed") + context_id = len(holder.contexts) + holder.contexts.append({ + "context_id": context_id, + "round": int(payload["round"]), + "scenario_id": int(source["scenario_id"]), + "gamma": float(source["gamma"]), + "step": int(source["step"]), + "state": np.asarray(source["state"], np.float32), + "hp10": np.asarray(source["hp10"], np.float32), + "low5": np.asarray(source["low5"], np.float32), + "hist": np.asarray(source["hist"], np.float32), + "ped_xy": np.asarray(source["ped_xy"], np.float32), + "ped_vel": np.asarray(source["ped_vel"], np.float32), + }) + row = { + "window_id": expected, + "query_id": expected, + "context_id": context_id, + "controls": np.asarray(source["controls"], np.float32), + "x0": np.asarray(source["x0"], np.float32), + "y": 0, + "semantic_label": "neutral", + } + holder.windows.append(row) + records.append((holder, row)) + if len(records) != int(payload["summary"]["D0"]): + raise RuntimeError("D0 payload count mismatch") + return holder, records + + +def _identity(holder, row): + return int(holder.round_i), int(row["query_id"]) + + +def _population_update( + policy, + optimizer, + records, + *, + population, + inner_steps, + batch, + device, + seed, +): + if not records: + raise ValueError(f"{population} replay requires nonempty support") + mass, accounting = BS.hierarchy_mass(records) + encoder_before = BS.module_sha256(policy.enc_grid) + expected = {_identity(holder, row) for holder, row in records} + losses = [] + encoder_gradient_norms = [] + exposure_hashes = [] + policy.train() + for inner in range(int(inner_steps)): + optimizer.zero_grad(set_to_none=True) + loss, visited = NS._objective( + policy, + records, + mass, + batch=int(batch), + device=device, + seed=int(seed) + inner * 1_000_003, + backward=True, + ) + identities = list(map(tuple, visited)) + if len(identities) != len(expected) or set(identities) != expected: + raise RuntimeError( + f"{population} pass duplicated or omitted a sample" + ) + squared = torch.zeros((), dtype=torch.float64) + for parameter in policy.enc_grid.parameters(): + if parameter.grad is not None: + squared += parameter.grad.detach().to( + dtype=torch.float64, + ).square().sum().cpu() + encoder_gradient_norms.append(float(squared.sqrt())) + optimizer.step() + losses.append(float(loss)) + exposure_hashes.append(_sha256_jsonable(identities)) + policy.eval() + encoder_after = BS.module_sha256(policy.enc_grid) + if ( + not any(parameter.requires_grad for parameter in policy.enc_grid.parameters()) + and encoder_after != encoder_before + ): + raise RuntimeError("visual encoder changed during replay") + return { + "population": str(population), + "records": len(records), + "inner_steps": int(inner_steps), + "optimizer_steps": int(inner_steps), + "sample_exposures": len(records) * int(inner_steps), + "exact_once_per_inner_step": True, + "exposure_identity_sha256": exposure_hashes, + "losses": losses, + "encoder_gradient_norms": encoder_gradient_norms, + "mass": _compact_mass(accounting), + "encoder_sha_before": encoder_before, + "encoder_sha_after": encoder_after, + } + + +def _skipped_population_update(records, *, population): + _, accounting = BS.hierarchy_mass(records) + return { + "population": str(population), + "records": len(records), + "inner_steps": 0, + "optimizer_steps": 0, + "sample_exposures": 0, + "exact_once_per_inner_step": None, + "exposure_identity_sha256": [], + "losses": [], + "encoder_gradient_norms": [], + "mass": _compact_mass(accounting), + "skipped": True, + "reason": "D0 collected for causal audit but neutral replay disabled", + } + + +@torch.no_grad() +def _encoder_probe(policy, records, *, device, limit=256): + """Return deterministic E_g tokens for a fixed record prefix.""" + ordered = sorted(records, key=lambda item: _identity(*item))[:int(limit)] + if not ordered: + raise ValueError("encoder probe requires records") + grid, _, _, _ = BS._tensor_batch(ordered, device) + was_training = policy.training + policy.eval() + token = policy.enc_grid(grid.float()).detach().cpu() + policy.train(was_training) + return token + + +def _encoder_probe_comparison(before, after): + before = torch.as_tensor(before, dtype=torch.float64).reshape(-1) + after = torch.as_tensor(after, dtype=torch.float64).reshape(-1) + if tuple(before.shape) != tuple(after.shape): + raise ValueError("encoder probe shapes changed") + cosine = float(torch.dot(before, after) / ( + before.norm() * after.norm() + ).clamp_min(1.0e-12)) + return { + "token_cosine": cosine, + "token_rms_change": float((after - before).square().mean().sqrt()), + "values": int(before.numel()), + } + + +@torch.no_grad() +def _encoder_probe_from_grid(policy, grid, *, device): + was_training = policy.training + policy.eval() + token = policy.enc_grid( + torch.as_tensor(grid, dtype=torch.float32, device=device) + ).detach().cpu() + policy.train(was_training) + return token + + +def _create_encoder_reference( + path, policy, records, *, device, anchor_round, checkpoint_sha256, + limit=256, +): + ordered = sorted(records, key=lambda item: _identity(*item))[:int(limit)] + if not ordered: + raise ValueError("encoder reference requires records") + grid, _, _, _ = BS._tensor_batch(ordered, device) + grid = grid.detach().cpu() + payload = { + "status": "SFM_B1_ENCODER_REFERENCE_PROBE", + "anchor_round": int(anchor_round), + "checkpoint_sha256": str(checkpoint_sha256), + "grid": grid, + "token": _encoder_probe_from_grid(policy, grid, device=device), + "encoder_snapshot": R2._module_snapshot(policy)["E_g"], + "records": len(ordered), + } + torch.save(payload, path) + return payload + + +def _load_encoder_reference(path): + payload = torch.load(path, map_location="cpu", weights_only=False) + if payload.get("status") != "SFM_B1_ENCODER_REFERENCE_PROBE": + raise RuntimeError("invalid encoder reference probe") + if not all(key in payload for key in ("grid", "token", "encoder_snapshot")): + raise RuntimeError("incomplete encoder reference probe") + return payload + + +def _encoder_reference_comparison(policy, reference, *, device): + current_token = _encoder_probe_from_grid( + policy, reference["grid"], device=device, + ) + current_snapshot = R2._module_snapshot(policy)["E_g"] + return { + **_encoder_probe_comparison(reference["token"], current_token), + "relative_parameter_drift": R2._module_relative_drift( + {"E_g": reference["encoder_snapshot"]}, + {"E_g": current_snapshot}, + )["E_g"], + "anchor_round": int(reference["anchor_round"]), + "records": int(reference["records"]), + } + + +def _gradient_diagnostics( + policy, positive_records, neutral_records, *, batch, device, seed, +): + positive = NS._gradient( + policy, + positive_records, + batch=batch, + device=device, + seed=seed, + ) + neutral = NS._gradient( + policy, + neutral_records, + batch=batch, + device=device, + seed=seed, + ) + return { + "Dplus_norm": BS._gradient_norm(positive), + "D0_norm": BS._gradient_norm(neutral), + "Dplus_D0_cosine": NS._cosine(positive, neutral), + } + + +def _anchor_catalog(trace_path, *, per_gamma, seed, output_path): + payload = torch.load( + trace_path, map_location="cpu", weights_only=False, + ) + candidates = [] + for trace in payload["traces"]: + if ( + trace.get("repair_trigger") is None + or trace.get("executed_controls") is None + or len(trace.get("all_K", ())) != 16 + or len(trace.get("selected_ids", ())) != 4 + ): + continue + all_x0 = np.stack([ + np.asarray(row["x0"], np.float32) + for row in trace["all_K"] + ]) + target_x0 = np.asarray(trace["executed_x0"], np.float32) + latent_distance = np.linalg.norm( + all_x0 - target_x0[None], axis=1, + ) + target_latent_id = int(np.argmin(latent_distance)) + if float(latent_distance[target_latent_id]) > 1.0e-7: + raise RuntimeError( + "executed target x0 is absent from the original K bank" + ) + base_lookup = { + int(row["candidate_id"]): row + for row in trace["query_rows"] + } + original_B_admissible = 0 + original_B_exact_positive = 0 + for candidate_id in map(int, trace["selected_ids"]): + row = base_lookup[candidate_id] + result = row["result"] + margin, _, _ = BC.nominal_hp_margin( + trace["state"], + row["controls"][0], + trace["ped_xy"], + trace["gamma"], + ) + original_B_exact_positive += int(result["y"]) + original_B_admissible += int( + int(result["y"]) == 1 and margin >= -1.0e-9 + ) + original_controls = np.stack([ + np.asarray(row["controls"], np.float32) + for row in trace["all_K"] + ]) + anchor = { + "anchor_id": len(candidates), + "round": int(trace["round"]), + "scenario_id": int(trace["scenario_id"]), + "gamma": float(trace["gamma"]), + "step": int(trace["step"]), + "trigger": str(trace["repair_trigger"]), + "population": ( + "D0" if trace["neutral_execution"] else "Dplus" + ), + "execution_source": str(trace["execution_source"]), + "state": np.asarray(trace["state"], np.float32), + "hp10": None, + "low5": None, + "hist": None, + "ped_xy": np.asarray(trace["ped_xy"], np.float32), + "ped_vel": np.asarray(trace["ped_vel"], np.float32), + "target_controls": np.asarray( + trace["executed_controls"], np.float32, + ), + "target_x0": target_x0, + "target_latent_id": target_latent_id, + "K_x0": all_x0, + "original_K_controls": original_controls, + "original_B_ids": list(map(int, trace["selected_ids"])), + "trace_B_exact_positive": original_B_exact_positive, + "trace_B_admissible": original_B_admissible, + "target_verifier_y": int( + trace["executed_result"]["y"] + ), + } + expected_y = 0 if anchor["population"] == "D0" else 1 + if anchor["target_verifier_y"] != expected_y: + raise RuntimeError( + "trigger target population disagrees with exact verifier y" + ) + candidates.append(anchor) + + # Context tensors are authoritative in the training stores, not rebuilt + # from floating-point state. + gather_dir = os.path.dirname(trace_path) + executed = OS.ExecutedRoundShard.load( + os.path.join(gather_dir, "executed_round.pt") + ) + neutral_payload = torch.load( + os.path.join(gather_dir, "neutral_round.pt"), + map_location="cpu", + weights_only=False, + ) + contexts = {} + for context in executed.contexts: + key = ( + int(context["scenario_id"]), + round(float(context["gamma"]), 8), + int(context["step"]), + ) + contexts[key] = context + for row in neutral_payload["records"]: + key = ( + int(row["scenario_id"]), + round(float(row["gamma"]), 8), + int(row["step"]), + ) + contexts[key] = row + for anchor in candidates: + key = ( + anchor["scenario_id"], + round(anchor["gamma"], 8), + anchor["step"], + ) + context = contexts.get(key) + if context is None: + raise RuntimeError(f"trigger anchor has no stored context: {key}") + anchor["hp10"] = np.asarray(context["hp10"], np.float32) + anchor["low5"] = np.asarray(context["low5"], np.float32) + anchor["hist"] = np.asarray(context["hist"], np.float32) + + grouped = defaultdict( + lambda: defaultdict(lambda: defaultdict(list)) + ) + for anchor in candidates: + grouped[round(anchor["gamma"], 8)][ + int(anchor["scenario_id"]) + ][anchor["population"]].append(anchor) + generator = np.random.default_rng(int(seed)) + selected = [] + for gamma in map(float, SP.GAMMAS): + key = round(gamma, 8) + values = [] + for scenario_id in sorted(grouped[key]): + populations = grouped[key][scenario_id] + for population in ("D0", "Dplus"): + rows = populations[population] + if rows: + index = int(generator.integers(0, len(rows))) + values.append(rows[index]) + values = values[:int(per_gamma)] + if len(values) < int(per_gamma): + used = {id(row) for row in values} + remainder = [ + row + for scenario_id in sorted(grouped[key]) + for population in ("D0", "Dplus") + for row in grouped[key][scenario_id][population] + if id(row) not in used + ] + remainder = [ + remainder[index] + for index in generator.permutation(len(remainder)) + ] + values.extend(remainder[:int(per_gamma) - len(values)]) + selected.extend(values) + for anchor_id, anchor in enumerate(selected): + anchor["anchor_id"] = anchor_id + catalog = { + "status": "SFM_B1_TRIGGER_ANCHOR_CATALOG_COMPLETE", + "round": int(payload["round"]), + "requested_per_gamma": int(per_gamma), + "candidate_count": len(candidates), + "selected_count": len(selected), + "selected_per_gamma": { + str(gamma): sum( + round(row["gamma"], 8) == round(float(gamma), 8) + for row in selected + ) + for gamma in SP.GAMMAS + }, + "selected_per_scenario_gamma": { + f"{scenario_id}:{gamma}": sum( + row["scenario_id"] == int(scenario_id) + and round(row["gamma"], 8) == round(float(gamma), 8) + for row in selected + ) + for scenario_id in sorted({ + row["scenario_id"] for row in candidates + }) + for gamma in SP.GAMMAS + }, + "anchors": selected, + } + temporary = output_path + ".tmp" + torch.save(catalog, temporary) + os.replace(temporary, output_path) + _write_json(output_path + ".COMPLETE.json", { + "status": catalog["status"], + "round": catalog["round"], + "file": os.path.abspath(output_path), + "sha256": FA._sha256_file(output_path), + "candidate_count": catalog["candidate_count"], + "selected_count": catalog["selected_count"], + "selected_per_gamma": catalog["selected_per_gamma"], + "selected_per_scenario_gamma": ( + catalog["selected_per_scenario_gamma"] + ), + }) + return selected + + +def _goal_progress(state, controls): + segment = SM.rollout_positions(state, controls) + initial = float(np.linalg.norm(segment[0] - SS.GOAL)) + one = float(np.linalg.norm(segment[1] - SS.GOAL)) + final = float(np.linalg.norm(segment[-1] - SS.GOAL)) + return initial - one, initial - final + + +def _predicted_clearance(state, controls, ped_xy, ped_vel): + robot = SM.rollout_positions(state, controls) + peds = SM.predict_pedestrians(ped_xy, ped_vel, len(controls)) + if peds.shape[1] == 0: + return float("inf") + return float( + np.linalg.norm(robot[:, None] - peds, axis=2).min() - SS.R_PED + ) + + +def _probe_checkpoint( + checkpoint, + anchors, + *, + phase, + device, + executor, + batch, + selector="margin", +): + policy, _ = GPS.load_sfm_policy(checkpoint, device=device) + policy.eval() + generated = [] + features = [] + reconstruction = [] + for start in range(0, len(anchors), int(batch)): + values = anchors[start:start + int(batch)] + hp10 = torch.as_tensor( + np.stack([row["hp10"] for row in values]), device=device, + ).float() + low5 = torch.as_tensor( + np.stack([row["low5"] for row in values]), device=device, + ).float() + hist = torch.as_tensor( + np.stack([row["hist"] for row in values]), device=device, + ).float() + context = policy.ctx_from(hp10, low5, hist) + x0 = torch.as_tensor( + np.stack([row["K_x0"] for row in values]), device=device, + ).float() + with torch.no_grad(): + controls = BE.integrate_latents( + policy, + x0.reshape(-1, policy.d), + context.repeat_interleave(16, dim=0), + nfe=8, + ).reshape(len(values), 16, SP.H, 2) + target = torch.as_tensor( + np.stack([row["target_controls"] for row in values]), + device=device, + ).float() + target_x0 = torch.as_tensor( + np.stack([row["target_x0"] for row in values]), + device=device, + ).float() + phi = policy.phi_s_from_x0( + target, + context, + target_x0, + s=0.9, + ) + controls_np = controls.detach().cpu().numpy().astype(np.float32) + generated.extend(controls_np) + features.extend( + BR.l2_normalize(phi).detach().cpu().numpy().astype(np.float32) + ) + reconstruction.extend([ + float(np.sqrt(np.mean( + (controls_np[index] - values[index][ + "original_K_controls" + ]) ** 2 + ))) + for index in range(len(values)) + ]) + + tasks = [] + for anchor_index, (anchor, windows) in enumerate( + zip(anchors, generated) + ): + for candidate_id, controls in enumerate(windows): + tasks.append(( + anchor_index, + candidate_id, + anchor["state"], + controls, + anchor["ped_xy"], + anchor["ped_vel"], + anchor["gamma"], + )) + result_rows = list(executor.map(SM.verify_in_worker, tasks)) + by_anchor = defaultdict(dict) + for anchor_index, candidate_id, result in result_rows: + if not result.get("resolved"): + raise RuntimeError("paired trigger probe verifier error") + by_anchor[int(anchor_index)][int(candidate_id)] = result + + rows = [] + for anchor_index, (anchor, windows, feature) in enumerate( + zip(anchors, generated, features) + ): + candidates = [] + admissible_count = 0 + for candidate_id, controls in enumerate(windows): + result = by_anchor[anchor_index][candidate_id] + margin, hp_old, hp_new = BC.nominal_hp_margin( + anchor["state"], + controls[0], + anchor["ped_xy"], + anchor["gamma"], + ) + admissible = bool( + int(result["y"]) == 1 and margin >= -1.0e-9 + ) + admissible_count += int(admissible) + candidates.append({ + "candidate_id": candidate_id, + "controls": controls, + "result": result, + "hp_margin": float(margin), + "hp_old": float(hp_old), + "hp_new": float(hp_new), + "admissible": admissible, + }) + B_orig = [ + candidates[index] for index in anchor["original_B_ids"] + ] + selected = BC.select_admissible( + B_orig, + selector=selector, + state=anchor["state"], + ped_xy=anchor["ped_xy"], + ped_vel=anchor["ped_vel"], + gamma=anchor["gamma"], + ) + target = anchor["target_controls"] + difference = windows - target[None] + same_latent = difference[int(anchor["target_latent_id"])] + target_latent = candidates[int(anchor["target_latent_id"])] + if selected is None: + selected_id = None + one_progress = None + H_progress = None + selected_margin = None + selected_clearance = None + else: + selected_id = int(selected["candidate_id"]) + one_progress, H_progress = _goal_progress( + anchor["state"], selected["controls"], + ) + selected_margin = float(selected["hp_margin"]) + selected_clearance = _predicted_clearance( + anchor["state"], + selected["controls"], + anchor["ped_xy"], + anchor["ped_vel"], + ) + rows.append({ + "anchor_id": int(anchor["anchor_id"]), + "round": int(anchor["round"]), + "scenario_id": int(anchor["scenario_id"]), + "gamma": float(anchor["gamma"]), + "step": int(anchor["step"]), + "trigger": str(anchor["trigger"]), + "population": str(anchor["population"]), + "execution_source": str(anchor["execution_source"]), + "phase": str(phase), + "target_verifier_y": int(anchor["target_verifier_y"]), + "target_full_rmse": float( + np.sqrt(np.mean(same_latent ** 2)) + ), + "target_first_rmse": float( + np.sqrt(np.mean(same_latent[0] ** 2)) + ), + "best_K_target_full_rmse": float( + np.sqrt(np.mean(difference ** 2, axis=(1, 2))).min() + ), + "baseline_K_reconstruction_rmse": float( + reconstruction[anchor_index] + ), + "all_K_exact_positive": int(sum( + int(row["result"]["y"]) for row in candidates + )), + "all_K_admissible": int(admissible_count), + "B_orig_admissible": int(sum( + row["admissible"] for row in B_orig + )), + "B_orig_NVP": selected is None, + "trace_B_exact_positive": int( + anchor["trace_B_exact_positive"] + ), + "trace_B_admissible": int(anchor["trace_B_admissible"]), + "baseline_B_admissible_delta_from_trace": int( + sum(row["admissible"] for row in B_orig) + - anchor["trace_B_admissible"] + ), + "target_latent_exact_y": int( + target_latent["result"]["y"] + ), + "target_latent_admissible": bool( + target_latent["admissible"] + ), + "selected_candidate_id": selected_id, + "selected_one_step_progress": one_progress, + "selected_H10_progress": H_progress, + "selected_hp_margin": selected_margin, + "selected_predicted_clearance": selected_clearance, + "_feature": np.asarray(feature, np.float32), + "_windows": np.asarray(windows, np.float32), + }) + del policy + if str(device).startswith("cuda"): + torch.cuda.empty_cache() + return rows + + +def _mean(values): + values = [float(value) for value in values if value is not None] + return None if not values else float(np.mean(values)) + + +def _probe_comparison(before, after): + before_lookup = {int(row["anchor_id"]): row for row in before} + after_lookup = {int(row["anchor_id"]): row for row in after} + if set(before_lookup) != set(after_lookup): + raise RuntimeError("paired trigger probe anchor set changed") + rows = [] + for anchor_id in sorted(before_lookup): + old = before_lookup[anchor_id] + new = after_lookup[anchor_id] + old_nvp = bool(old["B_orig_NVP"]) + new_nvp = bool(new["B_orig_NVP"]) + old_target_admissible = bool(old["target_latent_admissible"]) + new_target_admissible = bool(new["target_latent_admissible"]) + progress = new["selected_one_step_progress"] + actual_repair = bool( + old["population"] == "D0" + and old["trigger"] == "finite_B_NVP" + and old_nvp + and not new_nvp + and progress is not None + and float(progress) > 0.0 + ) + safe_stall = bool( + old["population"] == "D0" + and old["trigger"] == "finite_B_NVP" + and old_nvp + and not new_nvp + and (progress is None or float(progress) <= 0.0) + ) + imitation_only = bool( + old["population"] == "D0" + and new["target_full_rmse"] < old["target_full_rmse"] + and new_nvp + ) + feature_cosine = float(np.dot( + old["_feature"], new["_feature"], + ) / ( + np.linalg.norm(old["_feature"]) + * np.linalg.norm(new["_feature"]) + + 1.0e-12 + )) + rows.append({ + "anchor_id": anchor_id, + "gamma": old["gamma"], + "trigger": old["trigger"], + "population": old["population"], + "old_NVP": old_nvp, + "new_NVP": new_nvp, + "D0_actual_repair": actual_repair, + "D0_safe_stall": safe_stall, + "D0_imitation_only": imitation_only, + "Dplus_retained": bool( + old["population"] == "Dplus" + and old_target_admissible + and new_target_admissible + ), + "Dplus_regressed": bool( + old["population"] == "Dplus" + and old_target_admissible + and not new_target_admissible + ), + "Dplus_target_latent_to_safe": bool( + old["population"] == "Dplus" + and not old_target_admissible + and new_target_admissible + ), + "delta_all_K_admissible_fraction": ( + new["all_K_admissible"] + - old["all_K_admissible"] + ) / 16.0, + "delta_B_orig_admissible": ( + new["B_orig_admissible"] + - old["B_orig_admissible"] + ), + "delta_target_full_rmse": ( + new["target_full_rmse"] + - old["target_full_rmse"] + ), + "delta_selected_one_step_progress": ( + None + if old["selected_one_step_progress"] is None + or new["selected_one_step_progress"] is None + else ( + new["selected_one_step_progress"] + - old["selected_one_step_progress"] + ) + ), + "target_phi_s_cosine": feature_cosine, + "generated_K_RMS_drift": float(np.sqrt(np.mean( + (new["_windows"] - old["_windows"]) ** 2 + ))), + }) + + def summarize(values): + D0 = [row for row in values if row["population"] == "D0"] + Dplus = [ + row for row in values if row["population"] == "Dplus" + ] + return { + "contexts": len(values), + "D0_contexts": len(D0), + "Dplus_contexts": len(Dplus), + "D0_actual_repairs": sum( + row["D0_actual_repair"] for row in D0 + ), + "D0_safe_stalls": sum( + row["D0_safe_stall"] for row in D0 + ), + "D0_imitation_only": sum( + row["D0_imitation_only"] for row in D0 + ), + "Dplus_retained": sum( + row["Dplus_retained"] for row in Dplus + ), + "Dplus_regressed": sum( + row["Dplus_regressed"] for row in Dplus + ), + "Dplus_target_latent_to_safe": sum( + row["Dplus_target_latent_to_safe"] for row in Dplus + ), + "old_B_orig_NVP_rate": _mean([ + row["old_NVP"] for row in values + ]), + "new_B_orig_NVP_rate": _mean([ + row["new_NVP"] for row in values + ]), + "mean_delta_all_K_admissible_fraction": _mean([ + row["delta_all_K_admissible_fraction"] + for row in values + ]), + "mean_delta_target_full_rmse": _mean([ + row["delta_target_full_rmse"] for row in values + ]), + "mean_target_phi_s_cosine": _mean([ + row["target_phi_s_cosine"] for row in values + ]), + "mean_generated_K_RMS_drift": _mean([ + row["generated_K_RMS_drift"] for row in values + ]), + } + + return { + "pooled": summarize(rows), + "per_gamma": { + str(gamma): summarize([ + row for row in rows + if round(row["gamma"], 8) == round(float(gamma), 8) + ]) + for gamma in SP.GAMMAS + }, + "per_trigger": { + trigger: summarize([ + row for row in rows if row["trigger"] == trigger + ]) + for trigger in sorted({row["trigger"] for row in rows}) + }, + "rows": rows, + } + + +def _strip_probe_rows(rows): + return [ + { + key: value + for key, value in row.items() + if not key.startswith("_") + } + for row in rows + ] + + +def _raw_evaluation( + arms, + output_dir, + *, + ep0, + M, + bank_role, + noise_seed, + device, + workers, + executor=None, +): + OE.M_PER_GAMMA = int(M) + probe, _ = GPS.load_sfm_policy(arms[0]["checkpoint"], device="cpu") + noise, noise_meta = OE._noise_bank( + ep0=int(ep0), d=int(probe.d), seed=int(noise_seed), + ) + del probe + cache_dir = os.path.join(output_dir, "cache") + os.makedirs(cache_dir, exist_ok=True) + records = [] + owns_executor = executor is None + if owns_executor: + context = mp.get_context("spawn") + executor = ProcessPoolExecutor( + max_workers=int(workers), mp_context=context, + ) + try: + for arm in arms: + cell = OE._evaluate_checkpoint( + arm["checkpoint"], + scene_profile="double_density_velocity_ood", + ep0=int(ep0), + noise=noise, + noise_meta=noise_meta, + device=device, + cache_dir=cache_dir, + executor=executor, + ) + records.append({ + "name": arm["name"], + "round": arm["round"], + "phase": arm["phase"], + "checkpoint": arm["checkpoint"], + "checkpoint_sha256": FA._sha256_file( + arm["checkpoint"] + ), + "cell": cell, + }) + finally: + if owns_executor: + executor.shutdown(wait=True) + report = { + "status": f"SFM_B1_RAW_M{int(M)}_COMPLETE", + "scene_profile": "double_density_velocity_ood", + "bank_role": str(bank_role), + "bank": { + "ep0": int(ep0), + "M_per_gamma": int(M), + "scenario_ids": list(range(int(ep0), int(ep0) + int(M))), + }, + "noise_bank": noise_meta, + "policy_semantics": ( + "raw temperature=1, NFE=8, no acquisition, verifier, repair, " + "guidance, or fallback" + ), + "records": records, + } + _write_json(os.path.join(output_dir, "RAW_EVALUATION.json"), report) + return report + + +def _save_checkpoint(policy, path, metadata): + BX._save_checkpoint(policy, path, metadata) + return { + "checkpoint": os.path.abspath(path), + "checkpoint_sha256": FA._sha256_file(path), + } + + +def _parse_eval_rounds(value, final_round): + rounds = tuple(sorted({ + int(value) + for value in str(value).split(",") + if str(value).strip() + })) + if ( + not rounds + or rounds[0] != 0 + or any(value < 0 or value > int(final_round) for value in rounds) + or int(final_round) not in rounds + ): + raise ValueError( + "eval_rounds must include 0 and the final round, with no " + "out-of-range entries" + ) + return rounds + + +def _validate_resume_config(previous_config, current_config): + compatibility_defaults = { + "encoder_lr_ratio": 0.0, + "neutral_replay": True, + } + compatibility_normalized = {} + for key, value in current_config.items(): + if key in {"name", "rounds"}: + continue + previous_value = previous_config.get(key) + if key not in previous_config and key in compatibility_defaults: + previous_value = compatibility_defaults[key] + compatibility_normalized[key] = previous_value + if previous_value != value: + raise RuntimeError(f"resume config changed: {key}") + return compatibility_normalized + + +def _round_record_ref(path, record): + marker = os.path.abspath(os.fspath(path)) + checkpoint = os.path.abspath(record["checkpoint"]) + round_i = int(record["round"]) + post_positive = os.path.join( + os.path.dirname(checkpoint), f"round_{round_i:02d}_post_positive.pt", + ) + if FA._sha256_file(checkpoint) != record["checkpoint_sha256"]: + raise RuntimeError("round checkpoint digest mismatch") + return { + "path": marker, + "sha256": FA._sha256_file(marker), + "round": round_i, + "post_D0": checkpoint, + "post_D0_sha256": record["checkpoint_sha256"], + "post_Dplus": post_positive, + "post_Dplus_sha256": FA._sha256_file(post_positive), + } + + +def _delivery_lineage_refs(delivery, delivery_path, seen=()): + """Authenticate current and recursively resumed round artifacts.""" + delivery_path = os.path.abspath(os.fspath(delivery_path)) + if delivery_path in seen: + raise RuntimeError("cyclic resume delivery chain") + current_paths = [ + os.path.abspath(path) for path in delivery.get("round_records", ()) + ] + frozen_current = list(delivery.get("round_record_refs", ())) + computed_current = [] + for path in current_paths: + with open(path) as stream: + record = json.load(stream) + if record.get("status") != ROUND_STATUS: + raise RuntimeError("invalid resume round marker") + computed_current.append(_round_record_ref(path, record)) + if frozen_current: + if frozen_current != computed_current: + raise RuntimeError("current resume round snapshot changed") + current_refs = frozen_current + else: + current_refs = computed_current + + prior_refs = [] + resume = delivery.get("resume") + if resume: + prior_path = os.path.abspath(resume["delivery"]) + payload = open(prior_path, "rb").read() + if hashlib.sha256(payload).hexdigest() != resume.get("delivery_sha256"): + raise RuntimeError("nested resume delivery digest mismatch") + prior_delivery = json.loads(payload) + expected = _delivery_lineage_refs( + prior_delivery, prior_path, seen=(*seen, delivery_path), + ) + snapshot = list(resume.get("round_record_refs", ())) + if not snapshot: + raise RuntimeError("nested resume lacks frozen prior-round refs") + if snapshot != expected: + raise RuntimeError("nested resume round snapshot changed") + prior_refs = snapshot + return [*prior_refs, *current_refs] + + +def _resume_artifacts(resume_root, cfg, *, source_sha, scenario_ep0): + root = os.path.abspath(os.fspath(resume_root)) + delivery_path = os.path.join(root, "DELIVERY_COMPLETE.json") + if not os.path.isfile(delivery_path): + raise FileNotFoundError(delivery_path) + with open(delivery_path) as stream: + delivery = json.load(stream) + if delivery.get("status") != STATUS: + raise RuntimeError("resume source is not a completed neutral run") + if delivery.get("checkpoint_sha256") != source_sha: + raise RuntimeError("resume source uses another pretrained checkpoint") + previous_config = delivery.get("config", {}) + current_config = asdict(cfg) + compatibility_normalized = _validate_resume_config( + previous_config, current_config, + ) + current_round_records = list(delivery.get("round_records", ())) + if not current_round_records: + raise RuntimeError("resume delivery has no round records") + round_record_refs = _delivery_lineage_refs(delivery, delivery_path) + records = [] + for ref in round_record_refs: + if FA._sha256_file(ref["path"]) != ref["sha256"]: + raise RuntimeError("resume round marker digest mismatch") + with open(ref["path"]) as stream: + record = json.load(stream) + if int(record["round"]) != int(ref["round"]): + raise RuntimeError("resume round marker index changed") + records.append(record) + rounds = [int(record["round"]) for record in records] + if rounds != list(range(1, max(rounds) + 1)): + raise RuntimeError("resume rounds are not contiguous from one") + resume_round = rounds[-1] + if resume_round >= int(cfg.rounds): + raise RuntimeError("resume source already reaches requested final round") + expected_scenarios = set(range( + int(scenario_ep0), int(scenario_ep0) + 2 * resume_round + )) + observed_scenarios = { + int(scenario) + for record in records for scenario in record.get("scenarios", ()) + } + if observed_scenarios != expected_scenarios: + raise RuntimeError("resume scenario schedule is incomplete or changed") + final = records[-1] + checkpoint = os.path.join( + root, "checkpoints", f"round_{resume_round:02d}.pt" + ) + optimizer = os.path.join( + root, "checkpoints", f"round_{resume_round:02d}_optimizers.pt" + ) + previous_executed = os.path.join( + root, "rounds", f"round_{resume_round:02d}", + "gather", "executed_round.pt", + ) + encoder_reference = delivery.get("encoder_reference_probe") + if encoder_reference is not None: + encoder_reference_path = os.path.abspath(encoder_reference["path"]) + if ( + not os.path.isfile(encoder_reference_path) + or FA._sha256_file(encoder_reference_path) + != encoder_reference.get("sha256") + ): + raise RuntimeError("resume encoder reference digest mismatch") + else: + encoder_reference_path = None + if FA._sha256_file(checkpoint) != final["checkpoint_sha256"]: + raise RuntimeError("resume checkpoint hash mismatch") + if FA._sha256_file(optimizer) != final["optimizer_state"]["sha256"]: + raise RuntimeError("resume optimizer hash mismatch") + expected_shard_sha = final["gather"]["executed_shard"]["sha256"] + if FA._sha256_file(previous_executed) != expected_shard_sha: + raise RuntimeError("resume GP-support shard hash mismatch") + executed = OS.ExecutedRoundShard.load(previous_executed) + positive = OS.positive_records(executed) + support = { + float(holder.contexts[int(row["context_id"])]["gamma"]) + for holder, row in positive + } + if support != set(map(float, SP.GAMMAS)): + raise RuntimeError("resume GP support does not cover all gammas") + return { + "root": root, + "delivery": delivery_path, + "delivery_sha256": FA._sha256_file(delivery_path), + "round_record_refs": round_record_refs, + "resume_round": resume_round, + "checkpoint": checkpoint, + "checkpoint_sha256": final["checkpoint_sha256"], + "optimizer": optimizer, + "optimizer_sha256": final["optimizer_state"]["sha256"], + "previous_executed": previous_executed, + "previous_executed_sha256": expected_shard_sha, + "encoder_reference": encoder_reference_path, + "visual_encoder_sha256": delivery["visual_encoder_sha256"], + "next_scenarios": [ + int(scenario_ep0) + 2 * resume_round, + int(scenario_ep0) + 2 * resume_round + 1, + ], + "legacy_config_defaults": compatibility_normalized, + } + + +def _restore_optimizer( + optimizer, path, parameters, *, resume_round, inner_steps, + parameter_names=None, + neutral_replay=True, +): + payload = torch.load(path, map_location="cpu", weights_only=False) + if int(payload.get("round", -1)) != int(resume_round): + raise RuntimeError("optimizer round mismatch") + state = payload.get("optimizer", {}) + groups = list(state.get("param_groups", ())) + target_groups = list(optimizer.param_groups) + if len(groups) != len(target_groups): + raise RuntimeError("optimizer hyperparameters changed") + for source, target in zip(groups, target_groups): + source_hyper = {key: value for key, value in source.items() if key != "params"} + target_hyper = {key: value for key, value in target.items() if key != "params"} + if source_hyper != target_hyper: + raise RuntimeError("optimizer hyperparameters changed") + saved_names = payload.get("parameter_names") + if saved_names is not None and list(saved_names) != list(parameter_names or ()): + raise RuntimeError("optimizer named parameter order changed") + ids = [ + parameter_id + for group in groups for parameter_id in group.get("params", ()) + ] + if len(ids) != len(parameters) or len(state.get("state", {})) != len(parameters): + raise RuntimeError("optimizer parameter support is incomplete") + updates_per_round = 1 + int(bool(neutral_replay)) + expected_step = updates_per_round * int(resume_round) * int(inner_steps) + for parameter_id, parameter in zip(ids, parameters): + values = state["state"].get(parameter_id) + if values is None: + raise RuntimeError("optimizer parameter state is missing") + step = int(torch.as_tensor(values["step"]).item()) + if step != expected_step: + raise RuntimeError("optimizer Adam step does not match global round") + for key in ("exp_avg", "exp_avg_sq"): + if tuple(values[key].shape) != tuple(parameter.shape): + raise RuntimeError("optimizer moment shape mismatch") + optimizer.load_state_dict(state) + return { + "round": int(resume_round), + "expected_adam_step": expected_step, + "updates_per_round": updates_per_round, + "parameters": len(parameters), + "parameter_names_authenticated": saved_names is not None, + "legacy_parameter_order_reconstructed": saved_names is None, + "param_group_hyperparameters": [ + {key: value for key, value in group.items() if key != "params"} + for group in groups + ], + } + + +def run(args): + eval_rounds = _parse_eval_rounds(args.eval_rounds, args.rounds) + locked_eval_values = ( + args.locked_eval_checkpoint, + args.locked_eval_round, + args.locked_eval_sha256, + ) + if any(value is not None for value in locked_eval_values) and not all( + value is not None for value in locked_eval_values + ): + raise ValueError("locked prior-best evaluation requires path, round, and SHA") + locked_eval = None + if all(value is not None for value in locked_eval_values): + locked_round = int(args.locked_eval_round) + locked_checkpoint = os.path.abspath(args.locked_eval_checkpoint) + locked_sha = str(args.locked_eval_sha256) + if locked_round not in eval_rounds or locked_round in (0, int(args.rounds)): + raise ValueError("locked prior-best round must be an interior eval round") + if FA._sha256_file(locked_checkpoint) != locked_sha: + raise RuntimeError("locked prior-best checkpoint SHA changed") + locked_eval = { + "round": locked_round, + "checkpoint": locked_checkpoint, + "checkpoint_sha256": locked_sha, + } + cfg = StudyConfig( + name=str(args.name), + rounds=int(args.rounds), + scenarios_per_round=2, + lr=float(args.lr), + inner_steps=int(args.inner_steps), + batch=128, + ell=float(args.ell), + gp_cap=int(args.gp_cap), + ess_target=float(args.ess_target), + selector=str(args.selector), + encoder_lr_ratio=float(args.encoder_lr_ratio), + neutral_replay=bool(args.neutral_replay), + sample_seed=int(args.sample_seed), + audit_seed=int(args.audit_seed), + train_seed=int(args.train_seed), + probe_seed=int(args.probe_seed), + ).validate() + output_root = os.path.abspath(args.output_root) + if os.path.exists(output_root): + raise FileExistsError(f"refusing to reuse output root: {output_root}") + source = FA._source() + if not source["tracked_worktree_clean"]: + raise RuntimeError("study requires a clean frozen worktree") + checkpoint = os.path.abspath(args.checkpoint) + source_sha = FA._sha256_file(checkpoint) + if source_sha != RA.EXPECTED_CHECKPOINT_SHA256: + raise RuntimeError("source pretrained checkpoint SHA changed") + os.makedirs(output_root) + checkpoints_dir = os.path.join(output_root, "checkpoints") + os.makedirs(checkpoints_dir) + rounds_dir = os.path.join(output_root, "rounds") + os.makedirs(rounds_dir) + + resume = ( + _resume_artifacts( + args.resume_run_root, + cfg, + source_sha=source_sha, + scenario_ep0=int(args.scenario_ep0), + ) + if args.resume_run_root else None + ) + if resume and int(resume["resume_round"]) not in eval_rounds: + raise ValueError("resumed evaluation must include the anchor round") + policy_checkpoint = resume["checkpoint"] if resume else checkpoint + policy, _ = GPS.load_sfm_policy(policy_checkpoint, device=args.device) + frozen = BS.configure_expansion_trainability(policy) + if cfg.encoder_lr_ratio > 0.0: + policy.enc_grid.requires_grad_(True) + initial_visual_sha = BS.module_sha256(policy.enc_grid) + if resume and initial_visual_sha != resume["visual_encoder_sha256"]: + raise RuntimeError("resume visual encoder hash mismatch") + trainable_names = [ + name for name, parameter in policy.named_parameters() + if parameter.requires_grad + ] + effective_frozen_names = [ + name for name, parameter in policy.named_parameters() + if not parameter.requires_grad + ] + encoder_parameter_ids = { + id(parameter) for parameter in policy.enc_grid.parameters() + } + main_parameters = [ + parameter for parameter in policy.parameters() + if parameter.requires_grad + and id(parameter) not in encoder_parameter_ids + ] + encoder_parameters = [ + parameter for parameter in policy.enc_grid.parameters() + if parameter.requires_grad + ] + parameters = [*main_parameters, *encoder_parameters] + name_by_parameter_id = { + id(parameter): name for name, parameter in policy.named_parameters() + } + optimizer_parameter_names = [ + name_by_parameter_id[id(parameter)] for parameter in parameters + ] + optimizer_groups = [{"params": main_parameters, "lr": cfg.lr}] + if encoder_parameters: + optimizer_groups.append({ + "params": encoder_parameters, + "lr": cfg.lr * cfg.encoder_lr_ratio, + }) + optimizer = torch.optim.Adam(optimizer_groups) + optimizer_restore = ( + _restore_optimizer( + optimizer, + resume["optimizer"], + parameters, + resume_round=resume["resume_round"], + inner_steps=cfg.inner_steps, + parameter_names=optimizer_parameter_names, + neutral_replay=cfg.neutral_replay, + ) + if resume else None + ) + round0_path = os.path.join(checkpoints_dir, "round_00.pt") + if not resume: + _save_checkpoint(policy, round0_path, { + "study": STATUS, + "round": 0, + "phase": "pretrained", + "source_checkpoint": checkpoint, + "source_sha256": source_sha, + "study_config": asdict(cfg), + }) + current_checkpoint = policy_checkpoint + encoder_reference_path = ( + resume["encoder_reference"] if resume else None + ) + encoder_reference = ( + _load_encoder_reference(encoder_reference_path) + if encoder_reference_path else None + ) + previous_executed_path = ( + resume["previous_executed"] if resume else None + ) + history = [] + milestone_arms = [{ + "name": "r0", + "round": 0, + "phase": "pretrained", + "checkpoint": checkpoint, + }] + start_round = int(resume["resume_round"]) if resume else 0 + if resume: + milestone_arms.append({ + "name": f"r{start_round}", + "round": start_round, + "phase": "resume_anchor_after_D0", + "checkpoint": current_checkpoint, + }) + if locked_eval is not None: + if any(item["round"] == locked_eval["round"] for item in milestone_arms): + raise ValueError("locked prior-best duplicates an existing milestone") + milestone_arms.append({ + "name": f"locked_r{locked_eval['round']}", + "round": locked_eval["round"], + "phase": "locked_prior_best", + "checkpoint": locked_eval["checkpoint"], + }) + + context = mp.get_context("spawn") + with ProcessPoolExecutor( + max_workers=int(args.workers), mp_context=context, + ) as probe_executor: + for round_i in range(start_round + 1, cfg.rounds + 1): + started = time.perf_counter() + round_dir = os.path.join(rounds_dir, f"round_{round_i:02d}") + os.makedirs(round_dir) + scenarios = tuple(range( + int(args.scenario_ep0) + (round_i - 1) * 2, + int(args.scenario_ep0) + round_i * 2, + )) + gather_dir = os.path.join(round_dir, "gather") + current_sha = FA._sha256_file(current_checkpoint) + RA.collect( + current_checkpoint, + scenarios=scenarios, + gammas=tuple(map(float, SP.GAMMAS)), + scene_profile=cfg.scene_profile, + selector=cfg.selector, + device=args.device, + verifier_workers=int(args.workers), + sample_seed=cfg.sample_seed, + audit_seed=cfg.audit_seed, + ell=cfg.ell, + neutral_continuation=True, + round_i=round_i, + expected_checkpoint_sha256=current_sha, + previous_executed_path=previous_executed_path, + gp_cap=cfg.gp_cap, + ess_target=cfg.ess_target, + verifier_executor=probe_executor, + T=cfg.T, + outdir=gather_dir, + ) + executed = OS.ExecutedRoundShard.load( + os.path.join(gather_dir, "executed_round.pt") + ) + _, neutral_records = _neutral_records( + os.path.join(gather_dir, "neutral_round.pt") + ) + positive_records = OS.positive_records(executed) + if executed.Dminus: + raise RuntimeError("ordinary executed D unexpectedly has D-") + if not positive_records or not neutral_records: + raise RuntimeError("round requires nonempty D+ and D0") + + anchor_path = os.path.join( + round_dir, "trigger_anchors.pt", + ) + anchors = _anchor_catalog( + os.path.join(gather_dir, "repair_trace.pt"), + per_gamma=int(args.probe_per_gamma), + seed=cfg.probe_seed + round_i, + output_path=anchor_path, + ) + if not anchors: + raise RuntimeError("round produced no guidance-trigger anchors") + + phase_before = _probe_checkpoint( + current_checkpoint, + anchors, + phase="before_update", + device=args.device, + executor=probe_executor, + batch=cfg.batch, + selector=cfg.selector, + ) + gradient = _gradient_diagnostics( + policy, + positive_records, + neutral_records, + batch=cfg.batch, + device=args.device, + seed=cfg.train_seed + round_i * 10_000_019, + ) + fixed_before = { + "Dplus": NS._fixed_loss( + policy, + positive_records, + batch=cfg.batch, + device=args.device, + seed=cfg.train_seed + round_i, + ), + "D0": NS._fixed_loss( + policy, + neutral_records, + batch=cfg.batch, + device=args.device, + seed=cfg.train_seed + round_i, + ), + } + encoder_probe_records = [ + *positive_records, *neutral_records, + ] + if encoder_reference is None: + encoder_reference_path = os.path.join( + output_root, "encoder_reference_probe.pt", + ) + encoder_reference = _create_encoder_reference( + encoder_reference_path, + policy, + encoder_probe_records, + device=args.device, + anchor_round=start_round, + checkpoint_sha256=current_sha, + ) + encoder_probe_before = _encoder_probe( + policy, encoder_probe_records, device=args.device, + ) + + before_parameters = R2._module_snapshot(policy) + positive_update = _population_update( + policy, + optimizer, + positive_records, + population="Dplus", + inner_steps=cfg.inner_steps, + batch=cfg.batch, + device=args.device, + seed=cfg.train_seed + round_i * 1_000_003, + ) + after_positive_parameters = R2._module_snapshot(policy) + post_positive_path = os.path.join( + checkpoints_dir, + f"round_{round_i:02d}_post_positive.pt", + ) + _save_checkpoint(policy, post_positive_path, { + "study": STATUS, + "round": round_i, + "phase": "post_Dplus", + "study_config": asdict(cfg), + "Dplus_records": len(positive_records), + "D0_used": False, + }) + phase_positive = _probe_checkpoint( + post_positive_path, + anchors, + phase="after_Dplus", + device=args.device, + executor=probe_executor, + batch=cfg.batch, + selector=cfg.selector, + ) + gradient_after_positive = _gradient_diagnostics( + policy, + positive_records, + neutral_records, + batch=cfg.batch, + device=args.device, + seed=cfg.train_seed + round_i * 10_000_019, + ) + fixed_after_positive = { + "Dplus": NS._fixed_loss( + policy, + positive_records, + batch=cfg.batch, + device=args.device, + seed=cfg.train_seed + round_i, + ), + "D0": NS._fixed_loss( + policy, + neutral_records, + batch=cfg.batch, + device=args.device, + seed=cfg.train_seed + round_i, + ), + } + + if cfg.neutral_replay: + neutral_update = _population_update( + policy, + optimizer, + neutral_records, + population="D0", + inner_steps=cfg.inner_steps, + batch=cfg.batch, + device=args.device, + seed=( + cfg.train_seed + round_i * 1_000_003 + + 500_000_000 + ), + ) + else: + neutral_update = _skipped_population_update( + neutral_records, population="D0", + ) + after_neutral_parameters = R2._module_snapshot(policy) + round_checkpoint = os.path.join( + checkpoints_dir, f"round_{round_i:02d}.pt", + ) + checkpoint_marker = _save_checkpoint(policy, round_checkpoint, { + "study": STATUS, + "round": round_i, + "phase": ( + "post_Dplus_then_D0" + if cfg.neutral_replay else "post_Dplus_D0_audit_only" + ), + "study_config": asdict(cfg), + "Dplus_records": len(positive_records), + "D0_records": len(neutral_records), + "D0_original_verifier_y": 0, + "D0_gp_eligible": False, + "D0_used": bool(cfg.neutral_replay), + }) + optimizer_path = os.path.join( + checkpoints_dir, f"round_{round_i:02d}_optimizers.pt", + ) + torch.save({ + "round": round_i, + "optimizer": optimizer.state_dict(), + "parameter_names": optimizer_parameter_names, + }, optimizer_path) + phase_neutral = _probe_checkpoint( + round_checkpoint, + anchors, + phase=( + "after_D0" if cfg.neutral_replay + else "after_Dplus_D0_audit_only" + ), + device=args.device, + executor=probe_executor, + batch=cfg.batch, + selector=cfg.selector, + ) + fixed_after_neutral = { + "Dplus": NS._fixed_loss( + policy, + positive_records, + batch=cfg.batch, + device=args.device, + seed=cfg.train_seed + round_i, + ), + "D0": NS._fixed_loss( + policy, + neutral_records, + batch=cfg.batch, + device=args.device, + seed=cfg.train_seed + round_i, + ), + } + encoder_probe_after = _encoder_probe( + policy, encoder_probe_records, device=args.device, + ) + encoder_cumulative = _encoder_reference_comparison( + policy, encoder_reference, device=args.device, + ) + if ( + cfg.encoder_lr_ratio == 0.0 + and BS.module_sha256(policy.enc_grid) != initial_visual_sha + ): + raise RuntimeError("visual encoder changed") + + paired_positive = _probe_comparison( + phase_before, phase_positive, + ) + paired_neutral_increment = _probe_comparison( + phase_positive, phase_neutral, + ) + paired_total = _probe_comparison( + phase_before, phase_neutral, + ) + probe_payload = { + "status": PROBE_STATUS, + "round": round_i, + "semantics": { + "contexts": ( + "same stored Markov context and pedestrian state" + ), + "latent_bank": "same original K=16 x0 rows", + "B_orig": "same original RBF-selected candidate IDs", + "verifier": SM.verifier_manifest(), + "selection": ( + "exact y=1 AND nominal-Hp gate, then " + f"{cfg.selector}" + ), + "claim": ( + "local generator correction/resubstitution; not " + "closed-loop generalization" + ), + }, + "phases": { + "before_update": _strip_probe_rows(phase_before), + "after_Dplus": _strip_probe_rows(phase_positive), + "after_D0": _strip_probe_rows(phase_neutral), + }, + "comparisons": { + "Dplus_increment": paired_positive, + "D0_increment": paired_neutral_increment, + "total": paired_total, + }, + } + _write_json( + os.path.join(round_dir, "PAIRED_TRIGGER_PROBE.json"), + probe_payload, + ) + + phase_arms = [ + { + "name": f"r{round_i}_before", + "round": round_i, + "phase": "before_update", + "checkpoint": current_checkpoint, + }, + { + "name": f"r{round_i}_post_Dplus", + "round": round_i, + "phase": "after_Dplus", + "checkpoint": post_positive_path, + }, + { + "name": f"r{round_i}_post_D0", + "round": round_i, + "phase": ( + "after_D0" if cfg.neutral_replay + else "after_Dplus_D0_audit_only" + ), + "checkpoint": round_checkpoint, + }, + ] + raw_m2_dir = os.path.join(round_dir, "same_lineage_raw_M2") + os.makedirs(raw_m2_dir) + raw_m2 = _raw_evaluation( + phase_arms, + raw_m2_dir, + ep0=scenarios[0], + M=2, + bank_role=( + "same two scenarios as this round's gather; " + "paired fit diagnostic only" + ), + noise_seed=int(args.noise_seed) + round_i, + device=args.device, + workers=int(args.workers), + executor=probe_executor, + ) + + gather_complete = json.load(open( + os.path.join(gather_dir, "COMPLETE.json") + )) + record = { + "status": ROUND_STATUS, + "round": round_i, + "scenarios": list(scenarios), + "lineages": len(scenarios) * len(SP.GAMMAS), + "gather": gather_complete, + "Dplus": len(positive_records), + "D0": len(neutral_records), + "gradient_before_update": gradient, + "gradient_after_Dplus_before_D0": ( + gradient_after_positive + ), + "fixed_loss": { + "before": fixed_before, + "after_Dplus": fixed_after_positive, + "after_D0": fixed_after_neutral, + }, + "updates": { + "Dplus": positive_update, + "D0": neutral_update, + "phase_order": ( + ["Dplus", "D0"] + if cfg.neutral_replay else ["Dplus"] + ), + "same_lr_and_inner_steps": bool(cfg.neutral_replay), + }, + "parameter_relative_drift": { + "Dplus_increment": R2._module_relative_drift( + before_parameters, after_positive_parameters, + ), + "D0_increment": R2._module_relative_drift( + after_positive_parameters, after_neutral_parameters, + ), + "total": R2._module_relative_drift( + before_parameters, after_neutral_parameters, + ), + }, + "encoder_diagnostics": { + **_encoder_probe_comparison( + encoder_probe_before, encoder_probe_after, + ), + "relative_parameter_drift": ( + R2._module_relative_drift( + before_parameters, after_neutral_parameters, + )["E_g"] + ), + "lr": ( + 0.0 if cfg.encoder_lr_ratio == 0.0 + else cfg.lr * cfg.encoder_lr_ratio + ), + "trainable": bool(cfg.encoder_lr_ratio > 0.0), + "cumulative_from_reference": encoder_cumulative, + }, + "paired_trigger_probe": { + "file": os.path.join( + round_dir, "PAIRED_TRIGGER_PROBE.json", + ), + "Dplus_increment": paired_positive["pooled"], + "D0_increment": paired_neutral_increment["pooled"], + "total": paired_total["pooled"], + }, + "same_lineage_raw_M2": { + "file": os.path.join( + raw_m2_dir, "RAW_EVALUATION.json", + ), + "records": [ + { + "name": row["name"], + "pooled": row["cell"]["summary"]["pooled"], + } + for row in raw_m2["records"] + ], + }, + **checkpoint_marker, + "optimizer_state": { + "path": optimizer_path, + "sha256": FA._sha256_file(optimizer_path), + "persistent_across_rounds": True, + "single_Adam_across_active_updates": True, + "D0_replay_enabled": bool(cfg.neutral_replay), + }, + "wall_seconds": time.perf_counter() - started, + } + _write_json( + os.path.join(round_dir, "ROUND_COMPLETE.json"), record, + ) + with open( + os.path.join(output_root, "metrics.jsonl"), "a", + ) as stream: + stream.write(json.dumps(record, allow_nan=False) + "\n") + print(json.dumps({ + "round": round_i, + "Dplus": len(positive_records), + "D0": len(neutral_records), + "D0_actual_repairs": ( + paired_neutral_increment["pooled"][ + "D0_actual_repairs" + ] + ), + "Dplus_regressed": ( + paired_positive["pooled"]["Dplus_regressed"] + ), + "wall_seconds": record["wall_seconds"], + }), flush=True) + history.append(record) + if round_i in eval_rounds: + milestone_arms.append({ + "name": f"r{round_i}", + "round": round_i, + "phase": ( + "after_D0" if cfg.neutral_replay + else "after_Dplus_D0_audit_only" + ), + "checkpoint": round_checkpoint, + }) + current_checkpoint = round_checkpoint + previous_executed_path = os.path.join( + gather_dir, "executed_round.pt", + ) + + disjoint_dir = os.path.join( + output_root, f"disjoint_raw_M{int(args.eval_M)}", + ) + os.makedirs(disjoint_dir) + disjoint = _raw_evaluation( + milestone_arms, + disjoint_dir, + ep0=int(args.eval_ep0), + M=int(args.eval_M), + bank_role="disjoint fixed CRN metric bank", + noise_seed=int(args.noise_seed), + device=args.device, + workers=int(args.workers), + ) + complete = { + "status": STATUS, + "source": source, + "checkpoint": checkpoint, + "checkpoint_sha256": source_sha, + "config": asdict(cfg), + "trainability_configure_initially_frozen": frozen, + "frozen_parameters": effective_frozen_names, + "initial_visual_encoder_sha256": initial_visual_sha, + "visual_encoder_sha256": BS.module_sha256(policy.enc_grid), + "encoder_reference_probe": { + "path": encoder_reference_path, + "sha256": FA._sha256_file(encoder_reference_path), + "anchor_round": int(encoder_reference["anchor_round"]), + "records": int(encoder_reference["records"]), + }, + "optimizer_groups": [ + { + "name": "trunk_head_low_history", + "lr": cfg.lr, + "parameters": len(main_parameters), + }, + *([{ + "name": "enc_grid", + "lr": cfg.lr * cfg.encoder_lr_ratio, + "parameters": len(encoder_parameters), + }] if encoder_parameters else []), + ], + "rounds": int(cfg.rounds), + "rounds_run_this_invocation": len(history), + "resume": resume, + "optimizer_restore": optimizer_restore, + "locked_prior_best": locked_eval, + "trainable_parameter_names": trainable_names, + "round_records": [ + os.path.join( + rounds_dir, + f"round_{record['round']:02d}", + "ROUND_COMPLETE.json", + ) + for record in history + ], + "round_record_refs": [ + _round_record_ref( + os.path.join( + rounds_dir, + f"round_{record['round']:02d}", + "ROUND_COMPLETE.json", + ), + record, + ) + for record in history + ], + "disjoint_raw_evaluation": { + "file": os.path.join( + disjoint_dir, "RAW_EVALUATION.json", + ), + "M_per_gamma": int(args.eval_M), + "ep0": int(args.eval_ep0), + "rounds": list(eval_rounds), + "records": [ + { + "name": row["name"], + "round": row["round"], + "checkpoint": row["cell"]["checkpoint"], + "checkpoint_sha256": row["cell"]["checkpoint_sha256"], + "pooled": row["cell"]["summary"]["pooled"], + } + for row in disjoint["records"] + ], + }, + "scientific_scope": { + "same_lineage_M2": "fit/behavior diagnostic only", + "fixed_trigger_probe": ( + "causal local generator audit at remembered contexts" + ), + f"disjoint_M{int(args.eval_M)}": ( + "small screening metric, not final confirmation" + ), + "D0": ( + "teacher action with immutable exact y=0; never a " + "certificate-positive claim" + ), + "ordinary_replay_window": ( + "W=1 current-round executed D+ only; this follows the " + "full-trajectory one-pass study, not the older query-W2 B1" + ), + }, + } + finished_source = FA._source() + if ( + not finished_source["tracked_worktree_clean"] + or finished_source["commit"] != source["commit"] + ): + raise RuntimeError("source worktree changed during neutral training") + _write_json( + os.path.join(output_root, "DELIVERY_COMPLETE.json"), complete, + ) + return complete + + +def build_parser(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--checkpoint", required=True) + parser.add_argument( + "--resume-run-root", + help=( + "completed neutral run to continue with model, Adam, global " + "round/seed schedule, and previous-round GP support intact" + ), + ) + parser.add_argument("--output-root", required=True) + parser.add_argument("--name", default="lr3em5_s04") + parser.add_argument("--rounds", type=int, default=2) + parser.add_argument("--scenario-ep0", type=int, default=DEFAULT_SCENARIO_EP0) + parser.add_argument("--eval-ep0", type=int, default=DEFAULT_EVAL_EP0) + parser.add_argument("--eval-M", type=int, default=20) + parser.add_argument( + "--eval-rounds", + default="0,1,2", + help=( + "comma-separated disjoint-M20 checkpoints; must include 0 " + "and the final round" + ), + ) + parser.add_argument("--locked-eval-checkpoint") + parser.add_argument("--locked-eval-round", type=int) + parser.add_argument("--locked-eval-sha256") + parser.add_argument("--lr", type=float, default=DEFAULT_LR) + parser.add_argument("--inner-steps", type=int, default=DEFAULT_INNER_STEPS) + parser.add_argument( + "--selector", + choices=("margin", "progress_gated_margin"), + default="margin", + ) + parser.add_argument( + "--encoder-lr-ratio", type=float, default=0.0, + help="0 keeps E_g frozen; the only creative sanity value is 0.1", + ) + parser.add_argument( + "--no-neutral-replay", dest="neutral_replay", + action="store_false", + help="collect and audit D0 but skip its optimizer update", + ) + parser.set_defaults(neutral_replay=True) + parser.add_argument("--ell", type=float, default=RA.DEFAULT_ELL) + parser.add_argument("--gp-cap", type=int, default=DEFAULT_GP_CAP) + parser.add_argument("--ess-target", type=float, default=0.5) + parser.add_argument("--probe-per-gamma", type=int, default=4) + parser.add_argument("--device", default="cuda") + parser.add_argument("--workers", type=int, default=64) + parser.add_argument("--sample-seed", type=int, default=700_000) + parser.add_argument("--audit-seed", type=int, default=2_026_073_0) + parser.add_argument("--train-seed", type=int, default=2_026_073_1) + parser.add_argument("--probe-seed", type=int, default=2_026_073_2) + parser.add_argument("--noise-seed", type=int, default=2_026_073_3) + return parser + + +def main(argv=None): + args = build_parser().parse_args(argv) + if int(args.eval_M) not in (10, 20): + raise ValueError("screening metric bank must be M=10 or M=20/gamma") + run(args) + print(os.path.join( + args.output_root, "DELIVERY_COMPLETE.json", + )) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/overnight_run_07_12_sfm/sfm_b1_neutral_teacher_sanity.py b/overnight_run_07_12_sfm/sfm_b1_neutral_teacher_sanity.py new file mode 100644 index 0000000..fdac603 --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_neutral_teacher_sanity.py @@ -0,0 +1,1179 @@ +"""Round-1 sanity study for isolated neutral-repair CFM replay. + +The study has one immutable gather shared by every arm: + +1. collect ordinary exact-positive executed windows in ``D+``; +2. collect guided, exact-negative executed repairs separately in ``D0``; +3. apply the usual alpha=0 positive-only B1 update once; +4. clone that checkpoint and apply a dedicated whole-D0 CFM objective. + +``D0`` is never relabeled, inserted into D+/D-, or exposed to the RBF GP. +This is an in-sample diagnostic on the same twenty scenario seeds used for +gathering, not a generalization or safety claim. +""" +from __future__ import annotations + +import argparse +from collections import defaultdict +from concurrent.futures import ProcessPoolExecutor +import csv +import json +import math +import multiprocessing as mp +import os + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import torch + +import _paths # noqa: F401 +import grid_policy_sfm as GPS +import sfm_b1_cost as BC +import sfm_b1_eval as BE +import sfm_b1_expand as BX +import sfm_b1_full_episode_audit as FA +import sfm_b1_kazuki_repair_audit as RA +import sfm_b1_offline_eval as OE +import sfm_b1_offline_store as OS +import sfm_b1_r2_alpha_replay as R2 +import sfm_b1_store as BS +import sfm_metrics2 as SM +import sfm_protocol as SP +import sfm_scene as SS + + +STATUS = "SFM_B1_NEUTRAL_TEACHER_SANITY_COMPLETE" +TRAINING_STATUS = "SFM_B1_NEUTRAL_TEACHER_TRAINING_COMPLETE" +PROBE_STATUS = "SFM_B1_NEUTRAL_TEACHER_PROBE_COMPLETE" +DEFAULT_EP0 = 250_000 +DEFAULT_M = 20 +DEFAULT_ORDINARY_LR = 1.0e-5 +DEFAULT_NEUTRAL_LRS = (3.0e-5, 1.0e-4) +DEFAULT_NEUTRAL_STEPS = (1, 4, 16) +DEFAULT_NOISE_SEED = 2_026_073_0 +DEFAULT_TRAIN_SEED = 2_026_073_1 +DEFAULT_PROBE_PER_GAMMA = 5 + + +class _SingleRecent: + """Minimal one-round view required by the original B1 replay.""" + + def __init__(self, shard): + self._shard = shard + self.window = 1 + + @property + def rounds(self): + return [self._shard] + + def positive_records(self): + return OS.positive_records(self._shard) + + def negative_records(self): + return OS.negative_records(self._shard) + + +class _NeutralHolder: + """Context holder compatible with the production hierarchy helpers.""" + + def __init__(self, round_i): + self.round_i = int(round_i) + self.contexts = [] + self.windows = [] + + +def _write_json(path, payload): + path = os.path.abspath(os.fspath(path)) + os.makedirs(os.path.dirname(path), exist_ok=True) + temporary = path + ".tmp" + with open(temporary, "w") as stream: + json.dump(payload, stream, indent=2, sort_keys=True, allow_nan=False) + os.replace(temporary, path) + + +def _compact_mass(accounting): + return { + "total": float(accounting["total"]), + "gamma": { + str(key): float(value) + for key, value in accounting["gamma"].items() + }, + "cells": len(accounting["cells"]), + "contexts": len(accounting["contexts"]), + } + + +def _neutral_records(payload): + if payload.get("status") != "SFM_B1_NEUTRAL_ROUND_COMPLETE": + raise ValueError("not an authenticated neutral-round payload") + holder = _NeutralHolder(payload["round"]) + records = [] + for expected, source in enumerate(payload["records"]): + verifier = source.get("verifier_result", {}) + if ( + int(source["neutral_id"]) != expected + or source["population"] != "D0" + or source["semantic_label"] != "neutral" + or int(source["verifier_y"]) != 0 + or not verifier.get("resolved") + or int(verifier.get("y", -1)) != 0 + or not verifier.get("full_h") + or int(verifier.get("terminal_step", -1)) != SP.H + or source["train_eligible"] + or source["replay_default"] + or source["gp_eligible"] + ): + raise RuntimeError("D0 semantics changed before teacher replay") + context_id = len(holder.contexts) + holder.contexts.append({ + "context_id": context_id, + "round": int(payload["round"]), + "scenario_id": int(source["scenario_id"]), + "gamma": float(source["gamma"]), + "step": int(source["step"]), + "state": np.asarray(source["state"], np.float32), + "hp10": np.asarray(source["hp10"], np.float32), + "low5": np.asarray(source["low5"], np.float32), + "hist": np.asarray(source["hist"], np.float32), + "ped_xy": np.asarray(source["ped_xy"], np.float32), + "ped_vel": np.asarray(source["ped_vel"], np.float32), + }) + row = { + "window_id": expected, + "query_id": expected, + "context_id": context_id, + "controls": np.asarray(source["controls"], np.float32), + "x0": np.asarray(source["x0"], np.float32), + "y": 0, + "semantic_label": "neutral", + "original_verifier_y": 0, + } + if ( + tuple(row["controls"].shape) != (SP.H, 2) + or tuple(row["x0"].shape) != (2 * SP.H,) + ): + raise RuntimeError("invalid D0 controls/x0 shape") + holder.windows.append(row) + records.append((holder, row)) + if len(records) != int(payload["summary"]["D0"]): + raise RuntimeError("D0 payload count mismatch") + return holder, records + + +def _objective( + policy, records, mass, *, batch, device, seed, backward, +): + ordered = BS.hierarchical_order(records, int(seed)) + visited = [] + total_loss = 0.0 + for start in range(0, len(ordered), int(batch)): + values = ordered[start:start + int(batch)] + grid, low, hist, controls = BS._tensor_batch(values, device) + context = policy.ctx_from(grid, low, hist) + weights = torch.as_tensor( + [ + len(values) * mass[ + (id(holder), int(row["query_id"])) + ] + for holder, row in values + ], + dtype=controls.dtype, + device=device, + ) + torch.manual_seed(int(seed) + start) + loss = policy.cfm_loss(controls, context, weights=weights) + if not bool(torch.isfinite(loss)): + raise FloatingPointError("non-finite whole-buffer CFM loss") + if backward: + loss.backward() + total_loss += float(loss.detach()) + visited.extend( + (int(holder.round_i), int(row["query_id"])) + for holder, row in values + ) + expected = { + (int(holder.round_i), int(row["query_id"])) + for holder, row in records + } + if len(visited) != len(expected) or set(visited) != expected: + raise RuntimeError("whole-buffer objective duplicated or omitted support") + return total_loss, visited + + +def _fixed_loss(policy, records, *, batch, device, seed): + mass, _ = BS.hierarchy_mass(records) + with torch.no_grad(): + loss, _ = _objective( + policy, + records, + mass, + batch=batch, + device=device, + seed=seed, + backward=False, + ) + return float(loss) + + +def _neutral_update( + policy, + optimizer, + records, + *, + steps, + batch, + device, + seed, +): + if not records or int(steps) < 1: + raise ValueError("neutral replay requires nonempty D0 and positive steps") + mass, accounting = BS.hierarchy_mass(records) + encoder_before = BS.module_sha256(policy.enc_grid) + losses = [] + policy.train() + for step in range(int(steps)): + optimizer.zero_grad(set_to_none=True) + loss, visited = _objective( + policy, + records, + mass, + batch=batch, + device=device, + seed=int(seed) + step * 1_000_003, + backward=True, + ) + optimizer.step() + if len(visited) != len(records): + raise RuntimeError("neutral exposure count mismatch") + losses.append(float(loss)) + policy.eval() + encoder_after = BS.module_sha256(policy.enc_grid) + if encoder_after != encoder_before: + raise RuntimeError("visual encoder changed during neutral replay") + return { + "semantic_label": "neutral_teacher_despite_exact_y0", + "records": len(records), + "inner_steps": int(steps), + "optimizer_steps": int(steps), + "sample_exposures": len(records) * int(steps), + "exact_once_per_inner_step": True, + "fresh_cfm_base_each_exposure": True, + "stored_x0_used_for_training": False, + "losses": losses, + "mass": _compact_mass(accounting), + "encoder_sha_before": encoder_before, + "encoder_sha_after": encoder_after, + } + + +def _gradient(policy, records, *, batch, device, seed): + mass, _ = BS.hierarchy_mass(records) + was_training = policy.training + try: + # cuDNN GRU backward is unavailable in eval mode. This diagnostic + # must use the same train-mode network semantics as both replay phases. + policy.train() + policy.zero_grad(set_to_none=True) + _objective( + policy, + records, + mass, + batch=batch, + device=device, + seed=seed, + backward=True, + ) + return BS._gradient_snapshot(policy) + finally: + policy.zero_grad(set_to_none=True) + policy.train(was_training) + + +def _cosine(first, second): + common = sorted(set(first) & set(second)) + numerator = sum( + float((first[key].double() * second[key].double()).sum()) + for key in common + if first[key] is not None and second[key] is not None + ) + first_norm = sum( + float(first[key].double().square().sum()) + for key in common if first[key] is not None + ) ** 0.5 + second_norm = sum( + float(second[key].double().square().sum()) + for key in common if second[key] is not None + ) ** 0.5 + return ( + numerator / (first_norm * second_norm) + if first_norm and second_norm else None + ) + + +def _arm_name(lr, steps): + lr_text = f"{float(lr):.0e}".replace("-", "m") + return f"neutral_lr{lr_text}_s{int(steps):02d}" + + +def _summarize_ordinary(value): + if value["path"] != "positive_only": + raise RuntimeError("alpha=0 ordinary phase left positive-only path") + return { + "path": value["path"], + "eligible": int(value["eligible"]), + "visited": len(value["visited"]), + "loss": float(value["loss"]), + "optimizer_steps": int(value["optimizer_steps"]), + "mass": _compact_mass(value["mass"]), + } + + +def _train( + checkpoint, + gather_dir, + output_dir, + *, + neutral_lrs, + neutral_steps, + ordinary_lr, + batch, + device, + seed, +): + executed = OS.ExecutedRoundShard.load( + os.path.join(gather_dir, "executed_round.pt") + ) + neutral_payload = torch.load( + os.path.join(gather_dir, "neutral_round.pt"), + map_location="cpu", + weights_only=False, + ) + neutral_holder, neutral_records = _neutral_records(neutral_payload) + if executed.Dminus: + raise RuntimeError("neutral collector unexpectedly populated ordinary D-") + if not executed.Dplus or not neutral_records: + raise RuntimeError("sanity study requires nonempty D+ and D0") + + os.makedirs(output_dir) + checkpoint_dir = os.path.join(output_dir, "checkpoints") + os.makedirs(checkpoint_dir) + policy, _ = GPS.load_sfm_policy(checkpoint, device=device) + frozen = BS.configure_expansion_trainability(policy) + encoder_sha = BS.module_sha256(policy.enc_grid) + positive_records = OS.positive_records(executed) + ordinary_before = R2._module_snapshot(policy) + ordinary_fixed_before = { + "Dplus": _fixed_loss( + policy, + positive_records, + batch=batch, + device=device, + seed=int(seed) + 20_000_000, + ), + "D0": _fixed_loss( + policy, + neutral_records, + batch=batch, + device=device, + seed=int(seed) + 20_000_001, + ), + } + optimizer = torch.optim.Adam( + [parameter for parameter in policy.parameters() + if parameter.requires_grad], + lr=float(ordinary_lr), + ) + ordinary = BS.signed_update( + policy, + optimizer, + _SingleRecent(executed), + alpha=0.0, + batch=int(batch), + device=device, + seed=int(seed), + ) + ordinary = _summarize_ordinary(ordinary) + if ordinary["optimizer_steps"] != 1: + raise RuntimeError("ordinary alpha=0 phase must use one Adam step") + policy.eval() + ordinary_fixed_after = { + "Dplus": _fixed_loss( + policy, + positive_records, + batch=batch, + device=device, + seed=int(seed) + 20_000_000, + ), + "D0": _fixed_loss( + policy, + neutral_records, + batch=batch, + device=device, + seed=int(seed) + 20_000_001, + ), + } + ordinary_drift = R2._module_relative_drift( + ordinary_before, R2._module_snapshot(policy) + ) + if BS.module_sha256(policy.enc_grid) != encoder_sha: + raise RuntimeError("visual encoder changed during ordinary replay") + ordinary_checkpoint = os.path.join( + checkpoint_dir, "post_positive.pt" + ) + BX._save_checkpoint(policy, ordinary_checkpoint, { + "study": STATUS, + "phase": "ordinary_Dplus", + "alpha": 0.0, + "lr": float(ordinary_lr), + "optimizer_steps": 1, + "neutral_used": False, + }) + + positive_gradient = _gradient( + policy, + positive_records, + batch=batch, + device=device, + seed=int(seed) + 10_000_000, + ) + neutral_gradient = _gradient( + policy, + neutral_records, + batch=batch, + device=device, + seed=int(seed) + 10_000_000, + ) + gradient_cosine = _cosine(positive_gradient, neutral_gradient) + policy.zero_grad(set_to_none=True) + del policy, optimizer + if str(device).startswith("cuda"): + torch.cuda.empty_cache() + + arms = [ + { + "name": "r0", + "phase": "pretrained", + "checkpoint": os.path.abspath(checkpoint), + "checkpoint_sha256": FA._sha256_file(checkpoint), + "neutral_lr": None, + "neutral_steps": 0, + }, + { + "name": "post_positive", + "phase": "ordinary_Dplus", + "checkpoint": ordinary_checkpoint, + "checkpoint_sha256": FA._sha256_file(ordinary_checkpoint), + "neutral_lr": None, + "neutral_steps": 0, + }, + ] + arm_logs = {} + for lr in neutral_lrs: + for steps in neutral_steps: + name = _arm_name(lr, steps) + policy, _ = GPS.load_sfm_policy( + ordinary_checkpoint, device=device + ) + BS.configure_expansion_trainability(policy) + before = R2._module_snapshot(policy) + fixed_before = { + "Dplus": _fixed_loss( + policy, + positive_records, + batch=batch, + device=device, + seed=int(seed) + 20_000_000, + ), + "D0": _fixed_loss( + policy, + neutral_records, + batch=batch, + device=device, + seed=int(seed) + 20_000_001, + ), + } + optimizer = torch.optim.Adam( + [parameter for parameter in policy.parameters() + if parameter.requires_grad], + lr=float(lr), + ) + update = _neutral_update( + policy, + optimizer, + neutral_records, + steps=int(steps), + batch=int(batch), + device=device, + seed=int(seed) + 30_000_000, + ) + fixed_after = { + "Dplus": _fixed_loss( + policy, + positive_records, + batch=batch, + device=device, + seed=int(seed) + 20_000_000, + ), + "D0": _fixed_loss( + policy, + neutral_records, + batch=batch, + device=device, + seed=int(seed) + 20_000_001, + ), + } + drift = R2._module_relative_drift( + before, R2._module_snapshot(policy) + ) + arm_checkpoint = os.path.join(checkpoint_dir, f"{name}.pt") + BX._save_checkpoint(policy, arm_checkpoint, { + "study": STATUS, + "phase": "neutral_D0_teacher", + "ordinary_checkpoint": ordinary_checkpoint, + "neutral_semantic_label": "neutral", + "original_verifier_y": 0, + "neutral_lr": float(lr), + "neutral_inner_steps": int(steps), + "neutral_gp_eligible": False, + }) + arm_logs[name] = { + "update": update, + "fixed_loss": { + "before": fixed_before, + "after": fixed_after, + }, + "module_relative_parameter_drift": drift, + } + arms.append({ + "name": name, + "phase": "neutral_D0_teacher", + "checkpoint": arm_checkpoint, + "checkpoint_sha256": FA._sha256_file(arm_checkpoint), + "neutral_lr": float(lr), + "neutral_steps": int(steps), + }) + del policy, optimizer + if str(device).startswith("cuda"): + torch.cuda.empty_cache() + + report = { + "status": TRAINING_STATUS, + "source": FA._source(), + "checkpoint": os.path.abspath(checkpoint), + "checkpoint_sha256": FA._sha256_file(checkpoint), + "gather_dir": os.path.abspath(gather_dir), + "executed_shard": { + "D": len(executed.D), + "Dplus": len(executed.Dplus), + "Dminus": len(executed.Dminus), + }, + "neutral_shard": { + "D0": len(neutral_records), + "all_original_verifier_y": 0, + "gp_eligible": 0, + "ordinary_replay_eligible": 0, + }, + "ordinary": { + "alpha": 0.0, + "lr": float(ordinary_lr), + **ordinary, + "fixed_loss": { + "before": ordinary_fixed_before, + "after": ordinary_fixed_after, + }, + "module_relative_parameter_drift": ordinary_drift, + }, + "Dplus_D0_gradient_cosine_at_post_positive": gradient_cosine, + "frozen_parameters": frozen, + "visual_encoder_sha256": encoder_sha, + "arms": arms, + "arm_logs": arm_logs, + "neutral_holder_round": neutral_holder.round_i, + } + _write_json(os.path.join(output_dir, "TRAINING_COMPLETE.json"), report) + return report, neutral_holder, neutral_records + + +def _probe_subset(holder, records, per_gamma, seed): + grouped = defaultdict(list) + for pair in records: + context = holder.contexts[int(pair[1]["context_id"])] + grouped[round(float(context["gamma"]), 8)].append(pair) + generator = np.random.default_rng(int(seed)) + selected = [] + for gamma in map(float, SP.GAMMAS): + values = grouped[round(gamma, 8)] + order = generator.permutation(len(values)) + selected.extend(values[index] for index in order[:int(per_gamma)]) + return selected + + +def _probe_one_arm( + arm, + probe_records, + *, + K, + batch, + device, + seed, + executor, +): + policy, _ = GPS.load_sfm_policy(arm["checkpoint"], device=device) + policy.eval() + tasks = [] + measurements = [] + for start in range(0, len(probe_records), int(batch)): + values = probe_records[start:start + int(batch)] + grid, low, hist, teacher = BS._tensor_batch(values, device) + context = policy.ctx_from(grid, low, hist) + latent_rows = [] + for holder, row in values: + stored = np.asarray(row["x0"], np.float32) + generator = np.random.default_rng(FA._keyed_seed( + int(seed), + int(holder.contexts[int(row["context_id"])]["scenario_id"]), + f"{float(holder.contexts[int(row['context_id'])]['gamma']):.8f}", + int(holder.contexts[int(row["context_id"])]["step"]), + int(row["query_id"]), + )) + extra = generator.standard_normal( + (int(K) - 1, int(policy.d)), dtype=np.float32 + ) + latent_rows.append(np.concatenate([stored[None], extra], axis=0)) + latents = torch.as_tensor( + np.stack(latent_rows), device=device, dtype=context.dtype + ) + with torch.no_grad(): + windows = BE.integrate_latents( + policy, + latents.reshape(-1, policy.d), + context.repeat_interleave(int(K), dim=0), + nfe=8, + ).reshape(len(values), int(K), SP.H, 2) + windows_np = windows.detach().cpu().numpy().astype(np.float32) + teacher_np = teacher.detach().cpu().numpy().astype(np.float32) + for local, (holder, row) in enumerate(values): + context_row = holder.contexts[int(row["context_id"])] + difference = windows_np[local] - teacher_np[local][None] + full_rmse = np.sqrt(np.mean(difference ** 2, axis=(1, 2))) + first_rmse = np.sqrt(np.mean( + difference[:, 0] ** 2, axis=1 + )) + candidate_rows = [] + for candidate_id, controls in enumerate(windows_np[local]): + task_id = len(tasks) + tasks.append(( + task_id, + candidate_id, + context_row["state"], + controls, + context_row["ped_xy"], + context_row["ped_vel"], + context_row["gamma"], + )) + candidate_rows.append({ + "task_id": task_id, + "candidate_id": candidate_id, + "controls": controls, + }) + measurements.append({ + "neutral_id": int(row["query_id"]), + "scenario_id": int(context_row["scenario_id"]), + "gamma": float(context_row["gamma"]), + "step": int(context_row["step"]), + "same_latent_full_rmse": float(full_rmse[0]), + "same_latent_first_rmse": float(first_rmse[0]), + "min_K_full_rmse": float(full_rmse.min()), + "min_K_first_rmse": float(first_rmse.min()), + "candidate_rows": candidate_rows, + "state": context_row["state"], + "ped_xy": context_row["ped_xy"], + }) + results = list(executor.map(SM.verify_in_worker, tasks)) + lookup = { + int(context_index): result + for context_index, _, result in results + } + for measurement in measurements: + positives = [] + progresses = [] + margins = [] + for candidate in measurement.pop("candidate_rows"): + result = lookup[int(candidate["task_id"])] + if not result.get("resolved"): + raise RuntimeError("hard-context probe verifier error") + positives.append(int(result["y"])) + state = np.asarray(measurement["state"], np.float32) + next_state = BE._step(state, candidate["controls"][0]) + progresses.append(float( + np.linalg.norm(state[:2] - SS.GOAL) + - np.linalg.norm(next_state[:2] - SS.GOAL) + )) + margin, _, _ = BC.nominal_hp_margin( + state, + candidate["controls"][0], + measurement["ped_xy"], + measurement["gamma"], + ) + margins.append(float(margin)) + measurement.update( + all_K_positive_fraction=float(np.mean(positives)), + any_positive_K=bool(any(positives)), + fixed_B4_any_positive=bool(any(positives[:4])), + same_latent_positive=bool(positives[0]), + mean_one_step_progress=float(np.mean(progresses)), + same_latent_one_step_progress=float(progresses[0]), + mean_nominal_hp_margin=float(np.mean(margins)), + same_latent_nominal_hp_margin=float(margins[0]), + ) + measurement.pop("state") + measurement.pop("ped_xy") + + def summarize(values): + if not values: + return { + "contexts": 0, + "all_K_positive_fraction": None, + "any_positive_K": None, + "fixed_B4_any_positive": None, + "fixed_B4_NVP": None, + "same_latent_positive": None, + "same_latent_full_rmse": None, + "same_latent_first_rmse": None, + "min_K_full_rmse": None, + "mean_one_step_progress": None, + "mean_nominal_hp_margin": None, + } + return { + "contexts": len(values), + "all_K_positive_fraction": float(np.mean([ + value["all_K_positive_fraction"] for value in values + ])), + "any_positive_K": float(np.mean([ + value["any_positive_K"] for value in values + ])), + "fixed_B4_any_positive": float(np.mean([ + value["fixed_B4_any_positive"] for value in values + ])), + "fixed_B4_NVP": float(1.0 - np.mean([ + value["fixed_B4_any_positive"] for value in values + ])), + "same_latent_positive": float(np.mean([ + value["same_latent_positive"] for value in values + ])), + "same_latent_full_rmse": float(np.mean([ + value["same_latent_full_rmse"] for value in values + ])), + "same_latent_first_rmse": float(np.mean([ + value["same_latent_first_rmse"] for value in values + ])), + "min_K_full_rmse": float(np.mean([ + value["min_K_full_rmse"] for value in values + ])), + "mean_one_step_progress": float(np.mean([ + value["mean_one_step_progress"] for value in values + ])), + "mean_nominal_hp_margin": float(np.mean([ + value["mean_nominal_hp_margin"] for value in values + ])), + } + + summary = { + "pooled": summarize(measurements), + "per_gamma": { + str(gamma): summarize([ + value for value in measurements + if float(value["gamma"]) == float(gamma) + ]) + for gamma in SP.GAMMAS + }, + } + del policy + if str(device).startswith("cuda"): + torch.cuda.empty_cache() + return summary, measurements + + +def _run_probe( + arms, + holder, + records, + output_dir, + *, + per_gamma, + K, + batch, + device, + workers, + seed, +): + probe_records = _probe_subset(holder, records, per_gamma, seed) + actual_per_gamma = { + str(gamma): sum( + float(holder.contexts[int(row["context_id"])]["gamma"]) + == float(gamma) + for _, row in probe_records + ) + for gamma in SP.GAMMAS + } + context = mp.get_context("spawn") + arm_results = [] + with ProcessPoolExecutor( + max_workers=int(workers), mp_context=context + ) as executor: + for arm in arms: + summary, rows = _probe_one_arm( + arm, + probe_records, + K=K, + batch=batch, + device=device, + seed=seed, + executor=executor, + ) + arm_results.append({ + "arm": arm["name"], + "checkpoint": arm["checkpoint"], + "checkpoint_sha256": arm["checkpoint_sha256"], + "summary": summary, + "rows": rows, + }) + report = { + "status": PROBE_STATUS, + "semantics": { + "population": "fixed gamma-balanced subset of exact-negative D0", + "stored_x0": "candidate 0 reuses the original flow base", + "other_latents": "15 deterministic Gaussian controls", + "fixed_B4": ( + "candidate 0 plus three deterministic controls; diagnostic " + "only, not RBF acquisition" + ), + "teacher_label": "exact verifier y=0 throughout", + }, + "K": int(K), + "requested_contexts_per_gamma": int(per_gamma), + "actual_contexts_per_gamma": actual_per_gamma, + "arms": arm_results, + } + _write_json(os.path.join(output_dir, "HARD_CONTEXT_PROBE.json"), report) + return report + + +def _run_raw_evaluation( + arms, + output_dir, + *, + ep0, + M, + noise_seed, + device, + workers, +): + OE.M_PER_GAMMA = int(M) + probe, _ = GPS.load_sfm_policy(arms[0]["checkpoint"], device="cpu") + noise, noise_meta = OE._noise_bank( + ep0=int(ep0), d=int(probe.d), seed=int(noise_seed) + ) + del probe + cache_dir = os.path.join(output_dir, "cache") + os.makedirs(cache_dir) + records = [] + context = mp.get_context("spawn") + with ProcessPoolExecutor( + max_workers=int(workers), mp_context=context + ) as executor: + for arm in arms: + cell = OE._evaluate_checkpoint( + arm["checkpoint"], + scene_profile="double_density_velocity_ood", + ep0=int(ep0), + noise=noise, + noise_meta=noise_meta, + device=device, + cache_dir=cache_dir, + executor=executor, + ) + records.append({ + "arm": arm["name"], + "phase": arm["phase"], + "neutral_lr": arm["neutral_lr"], + "neutral_steps": arm["neutral_steps"], + "checkpoint": arm["checkpoint"], + "cell": cell, + }) + report = { + "status": f"SFM_B1_NEUTRAL_RAW_M{int(M)}_COMPLETE", + "scene_profile": "double_density_velocity_ood", + "bank": { + "ep0": int(ep0), + "M_per_gamma": int(M), + "scenario_ids": list(range(int(ep0), int(ep0) + int(M))), + "same_as_gathering_bank": True, + "interpretation": "paired in-sample/resubstitution sanity only", + }, + "noise_bank": noise_meta, + "policy_semantics": ( + "raw temperature=1, NFE=8, no acquisition, verifier, repair, " + "guidance, or fallback" + ), + "records": records, + } + _write_json(os.path.join(output_dir, "RAW_EVALUATION.json"), report) + return report + + +def _metric(cell, name): + if name in ("CR", "SR", "timeout"): + return float(cell[name]) + key = { + "Validity": "Validity", + "clearance": "successful_clearance", + "time": "successful_time_to_goal", + }[name] + value = cell[key]["mean"] + return None if value is None else float(value) + + +def _render(raw, probe, output_dir): + metrics = ( + ("CR", "Collision rate"), + ("Validity", "Validity"), + ("clearance", "Successful min. clearance [m]"), + ("time", "Successful time-to-goal [s]"), + ) + colors = plt.get_cmap("turbo")( + np.linspace(0.05, 0.95, len(raw["records"])) + ) + figure, axes = plt.subplots(2, 2, figsize=(14, 10), squeeze=False) + for axis, (metric, title) in zip(axes.flat, metrics): + for color, record in zip(colors, raw["records"]): + values = [ + _metric( + record["cell"]["summary"]["per_gamma"][str(gamma)], + metric, + ) + for gamma in SP.GAMMAS + ] + axis.plot( + SP.GAMMAS, + [np.nan if value is None else value for value in values], + marker="o", + ms=3.5, + lw=1.3, + color=color, + label=record["arm"], + ) + axis.set_title(title) + axis.set_xlabel(r"$\gamma$") + axis.grid(alpha=0.25) + figure.legend( + *axes[0, 0].get_legend_handles_labels(), + loc="upper center", + ncol=4, + frameon=False, + ) + figure.tight_layout(rect=(0, 0, 1, 0.91)) + png = os.path.join(output_dir, "raw_gamma_trends.png") + pdf = os.path.join(output_dir, "raw_gamma_trends.pdf") + figure.savefig(png, dpi=220, bbox_inches="tight") + figure.savefig(pdf, bbox_inches="tight") + plt.close(figure) + + probe_lookup = { + row["arm"]: row["summary"]["pooled"] + for row in probe["arms"] + } + csv_path = os.path.join(output_dir, "pooled_summary.csv") + with open(csv_path, "w", newline="") as stream: + writer = csv.DictWriter(stream, fieldnames=( + "arm", "neutral_lr", "neutral_steps", "SR", "CR", "timeout", + "Validity", "clearance", "time", "D0_same_latent_RMSE", + "D0_any_positive_K16", "D0_fixed_B4_NVP", + )) + writer.writeheader() + for record in raw["records"]: + pooled = record["cell"]["summary"]["pooled"] + local = probe_lookup[record["arm"]] + writer.writerow({ + "arm": record["arm"], + "neutral_lr": record["neutral_lr"], + "neutral_steps": record["neutral_steps"], + "SR": _metric(pooled, "SR"), + "CR": _metric(pooled, "CR"), + "timeout": _metric(pooled, "timeout"), + "Validity": _metric(pooled, "Validity"), + "clearance": _metric(pooled, "clearance"), + "time": _metric(pooled, "time"), + "D0_same_latent_RMSE": local["same_latent_full_rmse"], + "D0_any_positive_K16": local["any_positive_K"], + "D0_fixed_B4_NVP": local["fixed_B4_NVP"], + }) + return [png, pdf, csv_path] + + +def run(args): + output_root = os.path.abspath(args.output_root) + if os.path.exists(output_root): + raise FileExistsError(f"refusing to reuse output root: {output_root}") + source = FA._source() + if not source["tracked_worktree_clean"]: + raise RuntimeError("sanity study requires a clean frozen worktree") + os.makedirs(output_root) + scenarios = tuple(range(int(args.ep0), int(args.ep0) + int(args.M))) + gather_dir = os.path.join(output_root, "gather") + RA.collect( + args.checkpoint, + scenarios=scenarios, + gammas=tuple(map(float, SP.GAMMAS)), + scene_profile="double_density_velocity_ood", + selector="margin", + device=args.device, + verifier_workers=int(args.workers), + sample_seed=int(args.sample_seed), + audit_seed=int(args.audit_seed), + ell=float(args.ell), + neutral_continuation=True, + T=SP.T, + outdir=gather_dir, + ) + training_dir = os.path.join(output_root, "training") + training, holder, neutral_records = _train( + args.checkpoint, + gather_dir, + training_dir, + neutral_lrs=tuple(map(float, args.neutral_lrs)), + neutral_steps=tuple(map(int, args.neutral_steps)), + ordinary_lr=float(args.ordinary_lr), + batch=int(args.batch), + device=args.device, + seed=int(args.train_seed), + ) + probe_dir = os.path.join(output_root, "probe") + os.makedirs(probe_dir) + probe = _run_probe( + training["arms"], + holder, + neutral_records, + probe_dir, + per_gamma=int(args.probe_per_gamma), + K=16, + batch=int(args.batch), + device=args.device, + workers=int(args.workers), + seed=int(args.probe_seed), + ) + evaluation_dir = os.path.join(output_root, "evaluation") + os.makedirs(evaluation_dir) + raw = _run_raw_evaluation( + training["arms"], + evaluation_dir, + ep0=int(args.ep0), + M=int(args.M), + noise_seed=int(args.noise_seed), + device=args.device, + workers=int(args.workers), + ) + outputs = _render(raw, probe, output_root) + complete = { + "status": STATUS, + "source": source, + "checkpoint": os.path.abspath(args.checkpoint), + "checkpoint_sha256": FA._sha256_file(args.checkpoint), + "scene_profile": "double_density_velocity_ood", + "gather_bank": { + "ep0": int(args.ep0), + "M": int(args.M), + "scenarios": list(scenarios), + "gammas": list(map(float, SP.GAMMAS)), + "lineages": len(scenarios) * len(SP.GAMMAS), + }, + "ordinary": { + "alpha": 0.0, + "lr": float(args.ordinary_lr), + "whole_Dplus_adam_steps": 1, + }, + "neutral_sweep": { + "lrs": list(map(float, args.neutral_lrs)), + "whole_D0_inner_steps": list(map(int, args.neutral_steps)), + "GP_eligible": False, + "original_verifier_y": 0, + }, + "evaluation": { + "same_scenarios_as_gather": True, + "M_per_gamma": int(args.M), + "raw_temperature": 1.0, + "claim_scope": "paired in-sample sanity; not promotion evidence", + }, + "training_complete": os.path.join( + training_dir, "TRAINING_COMPLETE.json" + ), + "probe_complete": os.path.join( + probe_dir, "HARD_CONTEXT_PROBE.json" + ), + "raw_evaluation": os.path.join( + evaluation_dir, "RAW_EVALUATION.json" + ), + "outputs": outputs, + } + _write_json(os.path.join(output_root, "DELIVERY_COMPLETE.json"), complete) + return complete + + +def build_parser(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--checkpoint", required=True) + parser.add_argument("--output-root", required=True) + parser.add_argument("--ep0", type=int, default=DEFAULT_EP0) + parser.add_argument("--M", type=int, default=DEFAULT_M) + parser.add_argument("--ordinary-lr", type=float, default=DEFAULT_ORDINARY_LR) + parser.add_argument( + "--neutral-lrs", + type=float, + nargs="+", + default=DEFAULT_NEUTRAL_LRS, + ) + parser.add_argument( + "--neutral-steps", + type=int, + nargs="+", + default=DEFAULT_NEUTRAL_STEPS, + ) + parser.add_argument("--batch", type=int, default=128) + parser.add_argument("--ell", type=float, default=RA.DEFAULT_ELL) + parser.add_argument("--device", default="cuda") + parser.add_argument("--workers", type=int, default=64) + parser.add_argument("--sample-seed", type=int, default=700_000) + parser.add_argument("--audit-seed", type=int, default=2_026_073_0) + parser.add_argument("--train-seed", type=int, default=DEFAULT_TRAIN_SEED) + parser.add_argument("--probe-seed", type=int, default=2_026_073_2) + parser.add_argument("--noise-seed", type=int, default=DEFAULT_NOISE_SEED) + parser.add_argument( + "--probe-per-gamma", + type=int, + default=DEFAULT_PROBE_PER_GAMMA, + ) + return parser + + +def main(argv=None): + args = build_parser().parse_args(argv) + if int(args.M) != 20: + raise ValueError("canonical sanity is pinned to M=20 scenarios") + if tuple(map(float, args.neutral_lrs)) != DEFAULT_NEUTRAL_LRS: + raise ValueError( + f"neutral LR grid must remain {DEFAULT_NEUTRAL_LRS}" + ) + if tuple(map(int, args.neutral_steps)) != DEFAULT_NEUTRAL_STEPS: + raise ValueError( + f"neutral inner-step grid must remain {DEFAULT_NEUTRAL_STEPS}" + ) + value = run(args) + print(os.path.join(args.output_root, "DELIVERY_COMPLETE.json")) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/overnight_run_07_12_sfm/sfm_b1_offline_18arm_compare.py b/overnight_run_07_12_sfm/sfm_b1_offline_18arm_compare.py new file mode 100644 index 0000000..f5bda2c --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_offline_18arm_compare.py @@ -0,0 +1,212 @@ +"""Paired four-metric comparison of margin and SafeMPPI-cost 9-arm sweeps.""" +from __future__ import annotations + +import argparse +import csv +import json +import os +from pathlib import Path + +import matplotlib +matplotlib.use("Agg") +import matplotlib.pyplot as plt + +import run_sfm_b1_offline_9arm as RUN + + +STATUS = "SFM_B1_OFFLINE_SELECTOR_COMPARISON_COMPLETE" +SELECTOR_ORDER = ("margin", "safemppi_cost", "balanced_rank") + + +def _read_json(path): + with open(path) as stream: + return json.load(stream) + + +def _load(root, selector): + root = Path(root).resolve() + delivery = _read_json(root / "DELIVERY_COMPLETE.json") + if delivery.get("status") != "SFM_B1_OFFLINE_9ARM_DELIVERY_COMPLETE": + raise ValueError(f"incomplete 9-arm delivery: {root}") + contract = dict(delivery["contract"]) + observed = contract.get("execution_selector", "margin") + if observed != selector: + raise ValueError( + f"expected {selector} sweep, observed {observed} at {root}" + ) + aggregate = _read_json( + root / "evaluation" / "aggregate" / "AGGREGATE_COMPLETE.json" + ) + rows = list(aggregate["rows"]) + if len(rows) != 9 * (RUN.ROUNDS + 1): + raise ValueError(f"expected 99 aggregate rows at {root}, got {len(rows)}") + for row in rows: + row["selector"] = selector + return root, delivery, contract, aggregate, rows + + +def _paired_r0(rows): + fields = ("SR", "CR", "timeout", "Validity", "clearance", "time_to_goal") + values = [ + tuple(row[field] for field in fields) + for row in rows if int(row["round"]) == 0 + ] + if not values or any(value != values[0] for value in values[1:]): + raise ValueError("all 18 arms must share an identical raw-M50 r0 cell") + return dict(zip(fields, values[0])) + + +def _plot(rows, output): + selectors = tuple( + selector for selector in SELECTOR_ORDER + if any(row["selector"] == selector for row in rows) + ) + combinations = [ + (float(alpha), int(exposure)) + for alpha in RUN.ALPHAS for exposure in RUN.EXPOSURE_EPOCHS + ] + colors = plt.get_cmap("tab10") + color_for = { + combination: colors(index) + for index, combination in enumerate(combinations) + } + linestyles = { + "margin": "-", + "safemppi_cost": "--", + "balanced_rank": ":", + } + specs = ( + ("CR", "Collision rate", (-.03, 1.03)), + ("Validity", "Validity", (-.03, 1.03)), + ("clearance", "Min. clearance [m]", None), + ("time_to_goal", "Time-to-goal [s]", None), + ) + figure, axes = plt.subplots(2, 2, figsize=(15.5, 10.5)) + for axis, (key, title, ylim) in zip(axes.flat, specs): + for selector in selectors: + for alpha, exposure in combinations: + values = [ + row for row in rows + if row["selector"] == selector + and float(row["alpha"]) == alpha + and int(row["exposure_epochs"]) == exposure + ] + values.sort(key=lambda row: int(row["round"])) + axis.plot( + [row["round"] for row in values], + [ + float("nan") if row[key] is None else float(row[key]) + for row in values + ], + color=color_for[(alpha, exposure)], + linestyle=linestyles[selector], + linewidth=1.55, alpha=.9, + ) + axis.set_title(title) + axis.set_xlabel("Expansion round") + axis.set_xticks(range(RUN.ROUNDS + 1)) + axis.grid(alpha=.24) + if ylim is not None: + axis.set_ylim(*ylim) + handles = [ + plt.Line2D( + [0], [0], color=color_for[(alpha, exposure)], lw=2.6, + label=rf"$\alpha={alpha:g}$, exposure={exposure}", + ) + for alpha, exposure in combinations + ] + selector_labels = { + "margin": "max one-step margin", + "safemppi_cost": "native SafeMPPI cost", + "balanced_rank": "balanced safety + performance rank", + } + handles.extend( + plt.Line2D( + [0], [0], color="black", lw=2.4, + linestyle=linestyles[selector], label=selector_labels[selector], + ) + for selector in selectors + ) + figure.legend( + handles=handles, ncol=4, loc="upper center", + frameon=False, fontsize=8, + ) + figure.tight_layout(rect=(0, 0, 1, .89)) + artifacts = [] + for suffix in ("png", "pdf"): + path = output / f"paired_{9 * len(selectors)}arm_raw_m50.{suffix}" + figure.savefig(path, dpi=300, bbox_inches="tight") + artifacts.append(str(path.resolve())) + plt.close(figure) + return artifacts + + +def compare(margin_root, cost_root, output_dir, *, balanced_root=None): + output = Path(output_dir).resolve() + output.mkdir(parents=True, exist_ok=False) + loaded = [ + _load(margin_root, "margin"), + _load(cost_root, "safemppi_cost"), + ] + if balanced_root is not None: + loaded.append(_load(balanced_root, "balanced_rank")) + rows = [row for item in loaded for row in item[-1]] + r0 = _paired_r0(rows) + csv_path = output / f"paired_{9 * len(loaded)}arm_raw_m50.csv" + fields = ( + "selector", "arm", "alpha", "exposure_epochs", "round", + "SR", "CR", "timeout", "Validity", "clearance", "time_to_goal", + ) + with csv_path.open("w", newline="") as stream: + writer = csv.DictWriter(stream, fieldnames=fields) + writer.writeheader() + writer.writerows({key: row[key] for key in fields} for row in rows) + figures = _plot(rows, output) + best_by_selector = {} + selectors = tuple(item[-1][0]["selector"] for item in loaded) + for selector in selectors: + candidates = [ + row for row in rows + if row["selector"] == selector and int(row["round"]) > 0 + ] + best_by_selector[selector] = min(candidates, key=RUN._screening_key) + report = { + "status": STATUS, + "comparison_role": ( + "paired common-bank raw-M50 screening; selector is the only " + "factor added to the existing alpha x exposure grid" + ), + "margin_root": str(loaded[0][0]), + "safemppi_cost_root": str(loaded[1][0]), + "balanced_rank_root": ( + None if len(loaded) == 2 else str(loaded[2][0]) + ), + "paired_r0": r0, + "best_screening_cell_by_selector": best_by_selector, + "rows": len(rows), + "csv": str(csv_path.resolve()), + "figures": figures, + } + marker = output / "COMPARISON_COMPLETE.json" + temporary = marker.with_suffix(".json.tmp") + with temporary.open("w") as stream: + json.dump(report, stream, indent=2, allow_nan=False) + os.replace(temporary, marker) + return report + + +def main(argv=None): + parser = argparse.ArgumentParser() + parser.add_argument("--margin-root", required=True) + parser.add_argument("--safemppi-cost-root", required=True) + parser.add_argument("--balanced-rank-root") + parser.add_argument("--output-dir", required=True) + args = parser.parse_args(argv) + compare( + args.margin_root, args.safemppi_cost_root, args.output_dir, + balanced_root=args.balanced_rank_root, + ) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/sfm_b1_offline_eval.py b/overnight_run_07_12_sfm/sfm_b1_offline_eval.py new file mode 100644 index 0000000..48cad8f --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_offline_eval.py @@ -0,0 +1,926 @@ +"""Raw SFM evaluation with terminal-truncated executed-window Validity. + +Every checkpoint uses one fixed M/scenario/gamma seed and latent bank. The +controller is the unguided raw flow at one declared global sampling +temperature: it samples one H=10 plan per context and executes only its first +action. Acquisition, verifier selection, fallback, and guidance are absent. +The default temperature remains one; non-default values are intended for a +separate validation-selected, then locked, evaluation protocol. + +For an executed trajectory with ``N_tau`` controls, Validity is the mean of +the ``N_tau`` exact GREEN-verifier indicators. The window at start ``t`` uses +the actions actually executed from ``t`` onward and has +``H_t=min(10, N_tau-t)``. Terminal tails are neither dropped nor padded. +""" +from __future__ import annotations + +import argparse +from concurrent.futures import ProcessPoolExecutor +from dataclasses import dataclass, field +import hashlib +import json +import math +import multiprocessing as mp +import os +import re +from typing import Any + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import torch + +import _paths # noqa: F401 +import grid_feats as GF +import grid_policy_sfm as GPS +import sfm_b1_eval as BE +import sfm_hp_history as HH +import sfm_metrics2 as SM +import sfm_protocol as SP +import sfm_scene as SS + + +VERSION = "sfm_b1_offline_executed_window_v1" +DEFAULT_M_PER_GAMMA = 50 +M_PER_GAMMA = DEFAULT_M_PER_GAMMA +T = int(SP.T) +H = int(SP.H) +NFE = 8 +TEMPERATURE = 1.0 +TEMPERATURE_BY_GAMMA = tuple(1.0 for _ in SP.GAMMAS) +DEFAULT_EP0 = 260_000 +DEFAULT_NOISE_SEED = 2_026_072_3 +Z95 = 1.959963984540054 +PLOT_SPECS = ( + ("CR", "Collision rate", (-0.03, 1.03)), + ("Validity", "Validity", (-0.03, 1.03)), + ("clearance", "Min. clearance [m]", None), + ("time", "Time-to-goal [s]", None), +) + + +def _artifact_prefix() -> str: + return f"raw_m{M_PER_GAMMA}_offline" + + +def _status() -> str: + return f"SFM_B1_OFFLINE_RAW_M{M_PER_GAMMA}_COMPLETE" + + +def _sha256_file(path: str | os.PathLike[str]) -> str: + digest = hashlib.sha256() + with open(path, "rb") as stream: + for chunk in iter(lambda: stream.read(1 << 20), b""): + digest.update(chunk) + return digest.hexdigest() + + +def _sha256_json(payload: Any) -> str: + encoded = json.dumps( + payload, sort_keys=True, separators=(",", ":") + ).encode() + return hashlib.sha256(encoded).hexdigest() + + +def _write_json(path: str | os.PathLike[str], payload: Any) -> None: + path = os.path.abspath(os.fspath(path)) + os.makedirs(os.path.dirname(path), exist_ok=True) + temporary = path + ".tmp" + with open(temporary, "w") as stream: + json.dump(payload, stream, indent=2, allow_nan=False) + os.replace(temporary, path) + + +def _checkpoint_specs(checkpoints: list[str], labels: list[str]) -> list[dict]: + if len(checkpoints) != len(labels) or not checkpoints: + raise ValueError("--checkpoints and --labels need the same nonzero length") + if len(labels) != len(set(labels)): + raise ValueError("checkpoint labels must be unique") + specs = [] + for checkpoint, label in zip(checkpoints, labels): + match = re.fullmatch(r"r([0-9]+)", str(label)) + if match is None: + raise ValueError(f"checkpoint label {label!r} must have form r0, r1, ...") + path = os.path.abspath(checkpoint) + if not os.path.isfile(path): + raise FileNotFoundError(path) + specs.append(dict(label=str(label), round=int(match.group(1)), checkpoint=path)) + rounds = [spec["round"] for spec in specs] + if rounds != sorted(rounds) or len(rounds) != len(set(rounds)): + raise ValueError("checkpoint labels must be unique and increasing") + return specs + + +def _noise_bank(*, ep0: int, d: int, seed: int) -> tuple[np.ndarray, dict]: + contract = { + "version": VERSION, + "ep0": int(ep0), + "M_per_gamma": M_PER_GAMMA, + "gammas": list(map(float, SP.GAMMAS)), + "T": T, + "d": int(d), + "seed": int(seed), + "temperature": ( + TEMPERATURE + if len(set(TEMPERATURE_BY_GAMMA)) == 1 else None + ), + "temperature_by_gamma": list(TEMPERATURE_BY_GAMMA), + "NFE": NFE, + } + generator = np.random.default_rng(int(seed)) + values = generator.standard_normal( + (len(SP.GAMMAS), M_PER_GAMMA, T, int(d)), dtype=np.float32, + ) + metadata = { + **contract, + "dtype": "float32", + "shape": list(values.shape), + "sha256": hashlib.sha256(values.tobytes(order="C")).hexdigest(), + "CRN": ( + "same (gamma,scenario,step) latent across checkpoints; paired " + "scenario IDs across gamma" + ), + } + return values, metadata + + +def _temperature_description() -> str: + if len(set(TEMPERATURE_BY_GAMMA)) == 1: + return f"global {TEMPERATURE_BY_GAMMA[0]:g}" + return "by-gamma " + ",".join( + f"{gamma:g}:{temperature:g}" + for gamma, temperature in zip(SP.GAMMAS, TEMPERATURE_BY_GAMMA) + ) + + +def _latent_tensor( + latents: np.ndarray, + gamma_indices: list[int], + *, + device: torch.device | str, +) -> torch.Tensor: + values = torch.as_tensor(np.asarray(latents), device=device) + if len(values) != len(gamma_indices): + raise ValueError("one gamma index is required per latent") + if len(set(TEMPERATURE_BY_GAMMA)) == 1: + # Preserve the original global-temperature arithmetic exactly. Moving + # this multiply to NumPy changes borderline rollout outcomes. + return TEMPERATURE * values + scales = torch.as_tensor([ + TEMPERATURE_BY_GAMMA[int(index)] for index in gamma_indices + ], dtype=values.dtype, device=device) + return values * scales[:, None] + + +@dataclass +class _Episode: + gamma_index: int + rollout_index: int + episode: int + gamma: float + humans: list + state: np.ndarray = field(default_factory=lambda: np.zeros(4, np.float32)) + history: HH.HpHistory = field(default_factory=HH.HpHistory) + controls: list[np.ndarray] = field(default_factory=list) + states: list[np.ndarray] = field( + default_factory=lambda: [np.zeros(4, np.float32)] + ) + ped_xy: list[np.ndarray] = field(default_factory=list) + ped_vel: list[np.ndarray] = field(default_factory=list) + status: str | None = None + minimum_clearance: float = float("inf") + + +def _clearance(state: np.ndarray, ped_xy: np.ndarray) -> float: + if not len(ped_xy): + return float("inf") + return float( + np.linalg.norm(ped_xy - state[:2][None], axis=1).min() - SS.R_PED + ) + + +def _terminal_check(episode: _Episode, ped_xy: np.ndarray) -> bool: + clearance = _clearance(episode.state, ped_xy) + episode.minimum_clearance = min(episode.minimum_clearance, clearance) + if clearance < 0.0: + episode.status = "collision" + elif float(np.linalg.norm(episode.state[:2] - SS.GOAL)) < 0.5: + episode.status = "success" + return episode.status is not None + + +@torch.no_grad() +def run_batched_raw( + policy, + *, + scene_profile: str, + ep0: int, + noise: np.ndarray, + device: str, +) -> list[dict]: + """Evaluate all 7xM cells and retain the controls actually executed.""" + environment = SS.scene_profile(scene_profile) + expected = (len(SP.GAMMAS), M_PER_GAMMA, T, int(policy.d)) + if tuple(noise.shape) != expected or noise.dtype != np.float32: + raise ValueError(f"noise bank {noise.shape}/{noise.dtype} != {expected}/float32") + episodes = [ + _Episode( + gamma_index=gamma_index, + rollout_index=rollout_index, + episode=int(ep0) + rollout_index, + gamma=float(gamma), + humans=SS.make_humans( + int(ep0) + rollout_index, + 0, + environment["n_ped"], + tuple(environment["ped_speed_range"]), + ), + ) + for gamma_index, gamma in enumerate(SP.GAMMAS) + for rollout_index in range(M_PER_GAMMA) + ] + + for step in range(T): + active, hp10, lows, histories, latents = [], [], [], [], [] + for episode in episodes: + if episode.status is not None: + continue + ped_xy, ped_vel = SS.collect_humans(episode.humans) + if _terminal_check(episode, ped_xy): + continue + obstacles = np.concatenate([ + ped_xy, + np.full((len(ped_xy), 1), SS.R_PED, np.float32), + ], axis=1) + raw_grid = torch.as_tensor(GF.axis_grid( + episode.state[:2], + obstacles, + 0.0, + R=SS.R_SENSE, + sensing=SS.R_SENSE, + )) + active.append((episode, ped_xy.copy(), ped_vel.copy())) + hp10.append(episode.history.append(raw_grid)) + lows.append(torch.as_tensor( + GF.low5(episode.state, SS.GOAL, episode.gamma) + )) + histories.append(torch.as_tensor(GF.hist_pad( + np.asarray(episode.controls[-16:]) + if episode.controls else np.zeros((0, 2)), + 16, + ))) + latents.append(noise[ + episode.gamma_index, + episode.rollout_index, + step, + ]) + if not active: + break + + hp10_tensor = torch.stack(hp10).to(device) + low_tensor = torch.stack(lows).to(device) + history_tensor = torch.stack(histories).to(device) + context = policy.ctx_from(hp10_tensor, low_tensor, history_tensor) + scaled_latents = _latent_tensor( + np.asarray(latents), + [episode.gamma_index for episode, _, _ in active], + device=device, + ) + windows = BE.integrate_latents( + policy, + scaled_latents, + context, + nfe=NFE, + ).reshape(len(active), H, 2) + windows = windows.detach().cpu().numpy().astype(np.float32) + + for (episode, ped_xy, ped_vel), window in zip(active, windows): + if tuple(window.shape) != (H, 2): + raise RuntimeError(f"generated plan {window.shape} != {(H, 2)}") + action = window[0].copy() + episode.ped_xy.append(ped_xy) + episode.ped_vel.append(ped_vel) + episode.controls.append(action) + episode.state = BE._step(episode.state, action) + episode.states.append(episode.state.copy()) + SS.advance_humans(episode.humans, episode.state) + + rows = [] + for episode in episodes: + if episode.status is None: + ped_xy, _ = SS.collect_humans(episode.humans) + if not _terminal_check(episode, ped_xy): + episode.status = "timeout" + success = episode.status == "success" + rows.append({ + "episode": episode.episode, + "gamma": episode.gamma, + "status": episode.status, + "success": success, + "collision": episode.status == "collision", + "timeout": episode.status == "timeout", + "steps": len(episode.controls), + "time_to_goal": len(episode.controls) * SS.DT if success else None, + "min_clearance": float(episode.minimum_clearance), + "successful_clearance": ( + float(episode.minimum_clearance) if success else None + ), + "states": np.asarray(episode.states, np.float32), + "controls": np.asarray(episode.controls, np.float32), + "ped_xy": np.asarray(episode.ped_xy, np.float32), + "ped_vel": np.asarray(episode.ped_vel, np.float32), + }) + return rows + + +def _verify_executed_episode(row: dict) -> dict: + """Return the fractional GREEN validity of all executed window starts.""" + n_steps = int(row["steps"]) + states = np.asarray(row["states"], np.float32) + controls = np.asarray(row["controls"], np.float32) + ped_xy = np.asarray(row["ped_xy"], np.float32) + ped_vel = np.asarray(row["ped_vel"], np.float32) + expected = ( + len(states) == n_steps + 1 + and len(controls) == n_steps + and len(ped_xy) == n_steps + and len(ped_vel) == n_steps + ) + if not expected or (n_steps and tuple(controls.shape[1:]) != (2,)): + return { + "validity": 0.0, + "valid_windows": 0, + "evaluated_windows": 0, + "verifier_errors": 1, + } + if n_steps == 0: + return { + "validity": 0.0, + "valid_windows": 0, + "evaluated_windows": 0, + "verifier_errors": 0, + } + + valid_windows = 0 + for start in range(n_steps): + stop = min(start + H, n_steps) + result = SM.verify_executed_window( + states[start], + controls[start:stop], + ped_xy[start], + ped_vel[start], + float(row["gamma"]), + ) + if not result.get("resolved", False): + return { + "validity": valid_windows / n_steps, + "valid_windows": valid_windows, + "evaluated_windows": start, + "verifier_errors": 1, + } + if int(result.get("window_horizon", -1)) != stop - start: + return { + "validity": valid_windows / n_steps, + "valid_windows": valid_windows, + "evaluated_windows": start, + "verifier_errors": 1, + } + valid_windows += int(bool(result["y"])) + return { + "validity": valid_windows / n_steps, + "valid_windows": valid_windows, + "evaluated_windows": n_steps, + "verifier_errors": 0, + } + + +def _attach_validity(rows: list[dict], executor) -> list[dict]: + futures = [executor.submit(_verify_executed_episode, row) for row in rows] + compact = [] + omitted = {"states", "controls", "ped_xy", "ped_vel"} + for row, future in zip(rows, futures): + value = {key: item for key, item in row.items() if key not in omitted} + value.update(future.result()) + compact.append(value) + return compact + + +def _cluster_bootstrap_interval( + rows: list[dict], + key: str, + *, + seed: int, + draws: int = 2_000, +) -> list[float | None]: + episode_ids = sorted({int(row["episode"]) for row in rows}) + sums, counts = [], [] + for episode in episode_ids: + values = [ + row.get(key) for row in rows if int(row["episode"]) == episode + ] + finite = [ + float(value) for value in values + if value is not None and math.isfinite(float(value)) + ] + sums.append(sum(finite)) + counts.append(len(finite)) + if not episode_ids or not sum(counts): + return [None, None] + generator = np.random.default_rng(int(seed)) + indices = generator.integers( + 0, len(episode_ids), size=(int(draws), len(episode_ids)) + ) + numerator = np.asarray(sums, float)[indices].sum(axis=1) + denominator = np.asarray(counts, float)[indices].sum(axis=1) + samples = numerator[denominator > 0] / denominator[denominator > 0] + if not len(samples): + return [None, None] + return list(map(float, np.quantile(samples, [.025, .975]))) + + +def _summarize_one(rows: list[dict], seed: int) -> dict: + n = len(rows) + if n < 1: + raise ValueError("cannot summarize an empty cell") + successes = sum(bool(row["success"]) for row in rows) + collisions = sum(bool(row["collision"]) for row in rows) + timeouts = sum(bool(row["timeout"]) for row in rows) + if successes + collisions + timeouts != n: + raise RuntimeError("success, collision, and timeout must partition a cell") + validity = BE.bootstrap_mean( + [row["validity"] for row in rows], seed=seed + 2 + ) + valid_windows = sum(int(row["valid_windows"]) for row in rows) + evaluated_windows = sum(int(row["evaluated_windows"]) for row in rows) + validity.update( + valid_windows=valid_windows, + evaluated_windows=evaluated_windows, + window_weighted_fraction=( + valid_windows / evaluated_windows if evaluated_windows else 0.0 + ), + ) + return { + "n": n, + "SR": successes / n, + "SR_wilson95": BE.wilson(successes, n), + "CR": collisions / n, + "CR_wilson95": BE.wilson(collisions, n), + "timeout": timeouts / n, + "timeout_wilson95": BE.wilson(timeouts, n), + "Validity": validity, + "successful_clearance": BE.bootstrap_mean( + [row["successful_clearance"] for row in rows], seed=seed + ), + "successful_time_to_goal": BE.bootstrap_mean( + [row["time_to_goal"] for row in rows], seed=seed + 1 + ), + "verifier_errors": sum(int(row["verifier_errors"]) for row in rows), + } + + +def summarize(rows: list[dict], *, seed: int) -> dict: + per_gamma = { + str(gamma): _summarize_one( + [row for row in rows if float(row["gamma"]) == float(gamma)], + seed + index * 10, + ) + for index, gamma in enumerate(SP.GAMMAS) + } + pooled = _summarize_one(rows, seed + 100) + for metric, key in ( + ("SR", "success"), + ("CR", "collision"), + ("timeout", "timeout"), + ): + pooled[f"{metric}_cluster_bootstrap95"] = _cluster_bootstrap_interval( + rows, key, seed=seed + 200 + len(metric) + ) + pooled["Validity"]["cluster_bootstrap95"] = _cluster_bootstrap_interval( + rows, "validity", seed=seed + 208 + ) + pooled["successful_clearance"]["cluster_bootstrap95"] = ( + _cluster_bootstrap_interval( + rows, "successful_clearance", seed=seed + 300 + ) + ) + pooled["successful_time_to_goal"]["cluster_bootstrap95"] = ( + _cluster_bootstrap_interval( + rows, "time_to_goal", seed=seed + 301 + ) + ) + pooled["ci_method"] = ( + "scenario-cluster bootstrap across seven paired gamma rows" + ) + return {"pooled": pooled, "per_gamma": per_gamma} + + +def _assert_zero_verifier_errors(summary: dict) -> None: + cells = [summary["pooled"], *summary["per_gamma"].values()] + if any(int(cell["verifier_errors"]) != 0 for cell in cells): + raise RuntimeError("evaluation contains verifier errors") + + +def _cell_key( + *, + checkpoint_sha256: str, + scene_profile: str, + ep0: int, + noise_meta: dict, +) -> str: + return _sha256_json({ + "version": VERSION, + "evaluator_sha256": _sha256_file(__file__), + "checkpoint_sha256": checkpoint_sha256, + "scene_profile": scene_profile, + "ep0": int(ep0), + "M_per_gamma": M_PER_GAMMA, + "noise_bank": noise_meta, + "temperature": ( + TEMPERATURE + if len(set(TEMPERATURE_BY_GAMMA)) == 1 else None + ), + "temperature_by_gamma": list(TEMPERATURE_BY_GAMMA), + "NFE": NFE, + "T": T, + "H": H, + "validity": "executed sliding windows H_t=min(10,N_tau-t)", + "verifier": SM.verifier_manifest(), + }) + + +def _evaluate_checkpoint( + checkpoint: str, + *, + scene_profile: str, + ep0: int, + noise: np.ndarray, + noise_meta: dict, + device: str, + cache_dir: str, + executor, +) -> dict: + checkpoint_sha = _sha256_file(checkpoint) + key = _cell_key( + checkpoint_sha256=checkpoint_sha, + scene_profile=scene_profile, + ep0=ep0, + noise_meta=noise_meta, + ) + cache_path = os.path.join( + cache_dir, f"offline_cell_{checkpoint_sha[:12]}_{key[:12]}.json" + ) + if os.path.isfile(cache_path): + with open(cache_path) as stream: + payload = json.load(stream) + if ( + payload.get("status") != "SFM_B1_OFFLINE_RAW_CELL_COMPLETE" + or payload.get("cell_key") != key + ): + raise RuntimeError(f"stale evaluation cache: {cache_path}") + _assert_zero_verifier_errors(payload["summary"]) + return payload + + policy, _ = GPS.load_sfm_policy(checkpoint, device=device) + policy.eval() + if int(policy.d) != int(noise.shape[-1]): + raise ValueError("checkpoint latent dimension does not match the noise bank") + rows = run_batched_raw( + policy, + scene_profile=scene_profile, + ep0=ep0, + noise=noise, + device=device, + ) + del policy + if str(device).startswith("cuda"): + torch.cuda.empty_cache() + compact = _attach_validity(rows, executor) + summary = summarize( + compact, + seed=int(ep0) + int(checkpoint_sha[:8], 16) % 100_000, + ) + _assert_zero_verifier_errors(summary) + payload = { + "status": "SFM_B1_OFFLINE_RAW_CELL_COMPLETE", + "cell_key": key, + "checkpoint": os.path.abspath(checkpoint), + "checkpoint_sha256": checkpoint_sha, + "scene_profile": scene_profile, + "ep0": int(ep0), + "M_per_gamma": M_PER_GAMMA, + "summary": summary, + "rows": compact, + "metric_semantics": { + "policy": ( + "canonical unguided raw flow, temperature=" + f"{_temperature_description()}, " + "NFE=8, one " + "generated H=10 window per context, execute first action" + ), + "Validity": ( + "mean per-trajectory fraction of all executed window starts " + "whose terminal-truncated actual-action window is task-space " + "valid, collision-free, and exact GREEN-certified; " + "H_t=min(10,N_tau-t)" + ), + "clearance": ( + "mean of each successful trajectory's minimum pedestrian " + "clearance; failures are excluded" + ), + "time": "successful trajectories only", + "outcome_partition": "SR + CR + timeout = 1", + }, + } + _write_json(cache_path, payload) + return payload + + +def _metric_value(cell: dict, metric: str) -> float: + if metric == "CR": + return float(cell[metric]) + key = { + "Validity": "Validity", + "clearance": "successful_clearance", + "time": "successful_time_to_goal", + }[metric] + value = cell[key]["mean"] + return float("nan") if value is None else float(value) + + +def _metric_interval(cell: dict, metric: str, *, pooled: bool) -> list[float]: + if metric == "CR": + value = ( + cell["CR_cluster_bootstrap95"] + if pooled else cell["CR_wilson95"] + ) + else: + key = { + "Validity": "Validity", + "clearance": "successful_clearance", + "time": "successful_time_to_goal", + }[metric] + value = ( + cell[key]["cluster_bootstrap95"] + if pooled else cell[key]["interval95"] + ) + return [ + float("nan") if item is None else float(item) + for item in value + ] + + +def render(records: list[dict], output_dir: str) -> list[str]: + """Render the four metrics in the ball-evaluator paper style.""" + colors = { + gamma: plt.get_cmap("plasma")( + 0.08 + 0.84 * index / max(len(SP.GAMMAS) - 1, 1) + ) + for index, gamma in enumerate(SP.GAMMAS) + } + rounds = [int(record["round"]) for record in records] + plt.rcParams.update({ + "font.family": "serif", + "mathtext.fontset": "cm", + "font.serif": ["cmr10", "Computer Modern Roman", "DejaVu Serif"], + "axes.titlesize": 24, + "axes.labelsize": 20, + "xtick.labelsize": 17, + "ytick.labelsize": 17, + "legend.fontsize": 16, + "axes.unicode_minus": False, + "axes.formatter.use_mathtext": True, + }) + figure, axes = plt.subplots(2, 2, figsize=(14.6, 10.8), squeeze=False) + for axis, (metric, title, ylim) in zip(axes.flat, PLOT_SPECS): + for gamma in SP.GAMMAS: + cells = [ + record["cell"]["summary"]["per_gamma"][str(gamma)] + for record in records + ] + values = [_metric_value(cell, metric) for cell in cells] + intervals = [ + _metric_interval(cell, metric, pooled=False) for cell in cells + ] + axis.plot( + rounds, values, color=colors[gamma], lw=1.35, alpha=0.75 + ) + axis.fill_between( + rounds, + [value[0] for value in intervals], + [value[1] for value in intervals], + color=colors[gamma], + alpha=0.18, + linewidth=0, + ) + pooled = [ + record["cell"]["summary"]["pooled"] for record in records + ] + values = [_metric_value(cell, metric) for cell in pooled] + intervals = [ + _metric_interval(cell, metric, pooled=True) for cell in pooled + ] + axis.plot(rounds, values, color="black", lw=3.0) + axis.fill_between( + rounds, + [value[0] for value in intervals], + [value[1] for value in intervals], + color="black", + alpha=0.14, + linewidth=0, + ) + axis.set_title(title) + axis.grid(alpha=0.25) + axis.set_xlim(rounds[0] - 0.4, rounds[-1] + 0.4) + if ylim is not None: + axis.set_ylim(*ylim) + axis.set_xlabel("Expansion round") + + handles = [ + plt.Line2D( + [0], [0], color=colors[gamma], lw=2.2, + label=rf"$\gamma={gamma:g}$", + ) + for gamma in SP.GAMMAS + ] + handles.append( + plt.Line2D([0], [0], color="black", lw=3.0, label="pooled") + ) + figure.legend( + handles=handles, ncol=5, loc="upper center", frameon=False + ) + figure.tight_layout(rect=(0, 0, 1, 0.90)) + os.makedirs(output_dir, exist_ok=True) + outputs = [] + for suffix in ("png", "pdf"): + path = os.path.join( + output_dir, f"{_artifact_prefix()}_curves.{suffix}" + ) + figure.savefig(path, dpi=300, bbox_inches="tight") + outputs.append(path) + plt.close(figure) + manifest = os.path.join( + output_dir, f"{_artifact_prefix()}_curves.figure.json" + ) + _write_json(manifest, { + "status": "SFM_B1_OFFLINE_FIGURE_COMPLETE", + "rounds": rounds, + "gammas": list(map(float, SP.GAMMAS)), + "claim": ( + "fixed raw temperature " + f"{_temperature_description()} rollouts; Validity is the " + "trajectory-mean " + "fraction over every executed window start; terminal horizons use " + "H_t=min(10,N_tau-t); every indicator requires task-space bounds, " + "time-indexed collision avoidance, and the exact GREEN certificate" + ), + "confidence_bands": ( + "per-gamma Wilson for CR and trajectory bootstrap for continuous " + "metrics; pooled scenario-cluster bootstrap" + ), + "style_source": "safeMPPI_demo_3d/scripts/evaluate_ball_expansion.py", + }) + outputs.append(manifest) + return outputs + + +def run(args) -> dict: + global M_PER_GAMMA, TEMPERATURE, TEMPERATURE_BY_GAMMA + M_PER_GAMMA = int(args.m_per_gamma) + if M_PER_GAMMA <= 0: + raise ValueError("--m-per-gamma must be positive") + TEMPERATURE = float(args.temperature) + if not math.isfinite(TEMPERATURE) or TEMPERATURE <= 0.0: + raise ValueError("--temperature must be finite and positive") + schedule = ( + list(map(float, args.temperature_by_gamma)) + if args.temperature_by_gamma is not None + else [TEMPERATURE] * len(SP.GAMMAS) + ) + if ( + len(schedule) != len(SP.GAMMAS) + or any(not math.isfinite(value) or value <= 0.0 for value in schedule) + ): + raise ValueError( + "--temperature-by-gamma requires seven finite positive values" + ) + TEMPERATURE_BY_GAMMA = tuple(schedule) + if len(set(TEMPERATURE_BY_GAMMA)) == 1: + TEMPERATURE = TEMPERATURE_BY_GAMMA[0] + specs = _checkpoint_specs(args.checkpoints, args.labels) + output_dir = os.path.abspath(args.output_dir) + cache_dir = os.path.abspath( + args.cache_dir or os.path.join(output_dir, "cache") + ) + os.makedirs(output_dir, exist_ok=True) + os.makedirs(cache_dir, exist_ok=True) + + probe, _ = GPS.load_sfm_policy(specs[0]["checkpoint"], device="cpu") + noise, noise_meta = _noise_bank( + ep0=args.ep0, d=int(probe.d), seed=args.noise_seed + ) + del probe + records = [] + context = mp.get_context("spawn") + with ProcessPoolExecutor( + max_workers=int(args.workers), mp_context=context + ) as executor: + for spec in specs: + cell = _evaluate_checkpoint( + spec["checkpoint"], + scene_profile=args.scene_profile, + ep0=args.ep0, + noise=noise, + noise_meta=noise_meta, + device=args.device, + cache_dir=cache_dir, + executor=executor, + ) + records.append({ + "label": spec["label"], + "round": spec["round"], + "cell": cell, + }) + + outputs = render(records, output_dir) + result = { + "status": _status(), + "version": VERSION, + "scene_profile": args.scene_profile, + "environment": SS.scene_profile(args.scene_profile), + "bank": { + "ep0": int(args.ep0), + "M_per_gamma": M_PER_GAMMA, + "scenario_ids": list(range( + int(args.ep0), int(args.ep0) + M_PER_GAMMA + )), + "same_scenario_ids_for_every_gamma": True, + }, + "noise_bank": noise_meta, + "temperature": ( + TEMPERATURE + if len(set(TEMPERATURE_BY_GAMMA)) == 1 else None + ), + "temperature_by_gamma": list(TEMPERATURE_BY_GAMMA), + "records": records, + "outputs": outputs, + } + result_path = os.path.join( + output_dir, f"{_artifact_prefix()}_metrics.json" + ) + _write_json(result_path, result) + result["metrics_json"] = result_path + return result + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--checkpoints", nargs="+", required=True) + parser.add_argument("--labels", nargs="+", required=True) + parser.add_argument( + "--scene-profile", + default="double_density_velocity_ood", + choices=SS.SCIENTIFIC_EVAL_PROFILES, + ) + parser.add_argument("--ep0", type=int, default=DEFAULT_EP0) + parser.add_argument("--noise-seed", type=int, default=DEFAULT_NOISE_SEED) + parser.add_argument( + "--m-per-gamma", type=int, default=DEFAULT_M_PER_GAMMA + ) + parser.add_argument( + "--temperature", + type=float, + default=1.0, + help=( + "global raw-policy latent scale; select on a validation bank and " + "lock before any disjoint confirmation" + ), + ) + parser.add_argument( + "--temperature-by-gamma", + nargs="+", + type=float, + help=( + "seven temperatures ordered as the canonical gamma list; use " + "only after a schedule is locked on a separate calibration bank" + ), + ) + parser.add_argument("--device", default="cuda") + parser.add_argument("--workers", type=int, default=32) + parser.add_argument("--cache-dir") + parser.add_argument("--output-dir", required=True) + return parser + + +def main(argv=None) -> int: + args = build_parser().parse_args(argv) + result = run(args) + print(result["metrics_json"]) + for path in result["outputs"]: + print(path) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/overnight_run_07_12_sfm/sfm_b1_offline_exec.py b/overnight_run_07_12_sfm/sfm_b1_offline_exec.py new file mode 100644 index 0000000..2a72f5a --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_offline_exec.py @@ -0,0 +1,935 @@ +"""Offline executed-window SFM expansion. + +B=4 exact verifier queries are used only to search for the next action. One +H=10 plan per context enters D: the plan whose first action was executed. At a +finite-B NVP context, an independent raw temperature-one plan is exact-verified +and executed even when negative so that the offline simulator can continue. + +This is an offline data collector, not a certified deployment controller. +""" +from __future__ import annotations + +import argparse +from collections import Counter, defaultdict +from concurrent.futures import ProcessPoolExecutor +import copy +from dataclasses import asdict, dataclass +import hashlib +import json +import math +import os +import subprocess +import time + +import numpy as np +import torch + +import _paths # noqa: F401 +import grid_policy_sfm as GPS +import sfm_b1_cost as BC +import sfm_b1_eval as BE +import sfm_b1_expand as BX +import sfm_b1_full_episode_audit as FA +import sfm_b1_offline_replay as OR +import sfm_b1_offline_store as OS +import sfm_b1_rbf as BR +import sfm_b1_store as BS +import sfm_metrics2 as SM +import sfm_protocol as SP +import sfm_scene as SS + + +EXPECTED_CHECKPOINT_SHA256 = ( + "1b5179c935d3eeff8824967d707d64cc9bab273949ee1f0e4f190172bab1b215" +) +ELL_MULTIPLIER = 0.5 +CAP = 512 +GP_LAMBDA = 1.0e-2 +ALPHAS = (0.0, 0.01, 0.1) +EXPOSURE_EPOCHS = (1, 10, 100) +SCENE_PROFILE = "double_density_velocity_ood" +EXECUTION_SELECTORS = ("margin", "safemppi_cost", "balanced_rank") + + +@dataclass(frozen=True) +class OfflineConfig: + alpha: float + exposure_epochs: int + selector: str = "margin" + rounds: int = 10 + K: int = 16 + B: int = 4 + T: int = 180 + H: int = 10 + batch: int = 128 + lr: float = 1.0e-4 + ess_target: float = 0.5 + nfe: int = 8 + temp: float = 1.0 + phi_s: float = 0.9 + gp_lam: float = GP_LAMBDA + verifier_workers: int = 8 + seed: int = 20260724 + scene_profile: str = SCENE_PROFILE + smoke: bool = False + + def validate(self): + if self.selector not in EXECUTION_SELECTORS: + raise ValueError(f"selector must be one of {EXECUTION_SELECTORS}") + if float(self.alpha) not in ALPHAS: + raise ValueError(f"alpha must be one of {ALPHAS}") + if int(self.exposure_epochs) not in EXPOSURE_EPOCHS: + raise ValueError(f"exposure_epochs must be one of {EXPOSURE_EPOCHS}") + if ( + int(self.K), int(self.B), int(self.T), int(self.H), + int(self.batch), float(self.lr), float(self.ess_target), + float(self.gp_lam), float(self.temp), self.scene_profile, + ) != ( + 16, 4, 180, 10, 128, 1.0e-4, 0.5, + GP_LAMBDA, 1.0, SCENE_PROFILE, + ): + raise ValueError("offline executed-window scientific contract changed") + expected_rounds = 1 if self.smoke else 10 + if int(self.rounds) != expected_rounds: + raise ValueError( + f"rounds must be {expected_rounds} when smoke={self.smoke}" + ) + if int(self.verifier_workers) < 1: + raise ValueError("verifier_workers must be positive") + return self + + @property + def arm_name(self): + alpha = str(float(self.alpha)).replace(".", "p") + prefixes = { + "margin": "offline_exec", + "safemppi_cost": "offline_exec_safemppi_cost", + "balanced_rank": "offline_exec_balanced_rank", + } + prefix = prefixes[self.selector] + return ( + f"{prefix}_alpha{alpha}_" + f"exposures{int(self.exposure_epochs):03d}" + ) + + +def _write_json(path, payload): + path = os.path.abspath(os.fspath(path)) + os.makedirs(os.path.dirname(path), exist_ok=True) + temporary = path + ".tmp" + with open(temporary, "w") as stream: + json.dump(payload, stream, indent=2, allow_nan=False) + os.replace(temporary, path) + + +def _source(): + root = os.path.abspath(os.path.join(os.path.dirname(__file__), "..")) + commit = subprocess.check_output( + ["git", "rev-parse", "HEAD"], cwd=root, text=True, + ).strip() + dirty = bool(subprocess.check_output( + ["git", "status", "--porcelain"], cwd=root, text=True, + ).strip()) + return dict(commit=commit, tracked_worktree_clean=not dirty) + + +def _keyed_seed(base, *parts): + payload = json.dumps( + [int(base), *parts], separators=(",", ":"), sort_keys=False, + ).encode() + return int.from_bytes(hashlib.sha256(payload).digest()[:8], "little") % ( + 2 ** 63 - 1 + ) + + +@torch.no_grad() +def _keyed_windows( + policy, live, batch, *, K, round_i, step, source, seed, nfe, temp, +): + contexts = policy.ctx_from(batch["hp10"], batch["low"], batch["hist"]) + latent_parts = [] + for replica in live: + generator = np.random.default_rng(_keyed_seed( + seed, int(round_i), int(replica.scenario_id), + f"{float(replica.gamma):.8f}", int(step), str(source), + )) + latent_parts.append(generator.standard_normal( + (int(K), int(policy.d)), dtype=np.float32, + )) + x0 = torch.as_tensor( + np.stack(latent_parts), + device=contexts.device, + dtype=contexts.dtype, + ) + latents = x0 * float(temp) + expanded = contexts.repeat_interleave(int(K), dim=0) + windows = BE.integrate_latents( + policy, latents.reshape(-1, policy.d), expanded, nfe=int(nfe), + ) + return windows.reshape( + len(live), int(K), int(policy.H_pred), 2, + ), contexts, x0 + + +@torch.no_grad() +def _features_from_x0(phi_policy, windows, contexts, x0, s): + if windows.shape[:2] != x0.shape[:2]: + raise ValueError("window and x0 candidate axes disagree") + K = int(windows.shape[1]) + controls = windows.reshape(-1, windows.shape[-2], 2) + expanded_contexts = contexts.repeat_interleave(K, dim=0) + features = phi_policy.phi_s_from_x0( + controls, + expanded_contexts, + x0.reshape(-1, phi_policy.d), + s=float(s), + ) + return BR.l2_normalize(features).reshape(len(contexts), K, -1) + + +@torch.no_grad() +def _initial_lengthscale(policy, replicas, cfg, device): + """Mean pairwise distance of 50 balanced pretrained proposals.""" + live, batch = BX._stack_prepared(replicas, device) + windows, contexts, x0 = _keyed_windows( + policy, live, batch, K=1, round_i=0, step=-2, + source="ell_preflight", seed=cfg.seed, nfe=cfg.nfe, temp=cfg.temp, + ) + features = _features_from_x0( + policy, windows, contexts, x0, cfg.phi_s, + )[:, 0] + groups = { + float(gamma): [ + index for index, replica in enumerate(live) + if round(float(replica.gamma), 8) == round(float(gamma), 8) + ] + for gamma in SP.GAMMAS + } + selected = [ + index + for gamma in SP.GAMMAS + for index in groups[float(gamma)][:7] + ] + selected.append(groups[float(SP.GAMMAS[0])][7]) + if len(selected) != 50 or len(set(selected)) != 50: + raise RuntimeError("ell preflight requires 50 unique balanced proposals") + ell0 = BR.mean_pairwise_lengthscale(features[selected]) + return float(ell0), float(ell0 * ELL_MULTIPLIER), dict( + count=50, + balance="7 per gamma plus one extra gamma=0.1", + proposal_source="pretrained policy at round-1 initial OOD contexts", + representation="stored proposal x0 at s=0.9", + multiplier=ELL_MULTIPLIER, + ) + + +def _gamma_balanced_records(previous, *, cap, round_i, seed): + if previous is None: + return [], dict( + requested_cap=int(cap), selected=0, quota=int(cap) // len(SP.GAMMAS), + rotating_extra_gamma=None, per_gamma={str(gamma): 0 for gamma in SP.GAMMAS}, + shortfall={str(gamma): int(cap) // len(SP.GAMMAS) for gamma in SP.GAMMAS}, + ) + groups = {} + for gamma_index, gamma in enumerate(SP.GAMMAS): + records = [ + (previous, row) + for row in previous.Dplus + if round( + float(previous.contexts[int(row["context_id"])]["gamma"]), 8 + ) == round(float(gamma), 8) + ] + groups[float(gamma)] = BS.hierarchical_order( + records, int(seed) + gamma_index, + ) + quota = int(cap) // len(SP.GAMMAS) + rotation = (int(round_i) - 2) % len(SP.GAMMAS) + extra_gamma = float(SP.GAMMAS[rotation]) + required = { + float(gamma): quota + int(float(gamma) == extra_gamma) + for gamma in SP.GAMMAS + } + shortfall = { + str(gamma): max(0, required[float(gamma)] - len(groups[float(gamma)])) + for gamma in SP.GAMMAS + } + if any(shortfall.values()): + raise RuntimeError( + "strict gamma-balanced GP quota is unavailable; " + f"required={required}, shortfall={shortfall}" + ) + selected = [ + record + for gamma in SP.GAMMAS + for record in groups[float(gamma)][:required[float(gamma)]] + ] + per_gamma = Counter( + str(previous.contexts[int(row["context_id"])]["gamma"]) + for _, row in selected + ) + identities = [ + (int(shard.round_i), int(row["window_id"])) for shard, row in selected + ] + if len(identities) != len(set(identities)): + raise RuntimeError("GP buffer selection contains duplicates") + if len(selected) != int(cap): + raise RuntimeError( + f"expected {cap} GP records, selected {len(selected)}" + ) + return selected, dict( + requested_cap=int(cap), + selected=len(selected), + quota=quota, + rotating_extra_gamma=extra_gamma, + per_gamma={str(gamma): int(per_gamma[str(gamma)]) for gamma in SP.GAMMAS}, + shortfall=shortfall, + unique=True, + ) + + +@torch.no_grad() +def gp_from_previous( + phi_policy, previous, *, round_i, ell, cap, lam, phi_s, device, seed, +): + selected, selection = _gamma_balanced_records( + previous, cap=cap, round_i=round_i, seed=seed, + ) + gp = BR.RBFGP(float(ell), float(lam)) + if selected: + feature_parts = [] + for start in range(0, len(selected), 256): + hp10, low, hist, controls = BX._record_batch( + selected[start:start + 256], device, + ) + x0 = torch.as_tensor(np.stack([ + row["x0"] for _, row in selected[start:start + 256] + ]), device=device).float() + feature_parts.append(phi_policy.phi_s_from_x0( + controls, + phi_policy.ctx_from(hp10, low, hist), + x0, + s=float(phi_s), + )) + gp.set_buffer(torch.cat(feature_parts)) + identities = [ + (int(shard.round_i), int(row["window_id"])) for shard, row in selected + ] + return gp, identities, selection + + +@torch.no_grad() +def _calibrate_beta( + phi_policy, gp, replicas, cfg, device, *, round_i, +): + live, batch = BX._stack_prepared(replicas, device) + windows, contexts, x0 = _keyed_windows( + phi_policy, live, batch, K=cfg.K, round_i=round_i, step=-1, + source="beta_calibration", seed=cfg.seed, nfe=cfg.nfe, temp=cfg.temp, + ) + features = _features_from_x0( + phi_policy, windows, contexts, x0, cfg.phi_s, + ) + vectors = [] + for index, (replica, values) in enumerate(zip(live, features)): + generator = torch.Generator(device=values.device).manual_seed( + _keyed_seed( + cfg.seed, round_i, replica.scenario_id, + f"{replica.gamma:.8f}", "beta_order", + ) + ) + order = torch.randperm( + len(values), generator=generator, device=values.device, + ) + vectors.extend(gp.sequential_score_vectors(values, order, cfg.B)) + beta, ess = BR.solve_beta(vectors, target=cfg.ess_target) + return float(beta), float(ess) + + +def _finalize_alive(replicas): + for replica in replicas: + if not replica.alive: + continue + terminal_xy, _ = SS.collect_humans(replica.humans) + clearance = float( + np.linalg.norm( + terminal_xy - replica.state[:2][None], axis=1, + ).min() - SS.R_PED + ) + replica.minimum_clearance = min(replica.minimum_clearance, clearance) + if clearance < 0.0: + replica.status = "collision" + elif float(np.linalg.norm(replica.state[:2] - SS.GOAL)) < 0.5: + replica.status = "success" + else: + replica.status = "timeout" + replica.alive = False + + +def gather_offline_round( + policy, phi_policy, gp, beta, replicas, cfg, shard, device, executor, + *, round_i, +): + timers = Counter() + counts = Counter() + sigma_all, sigma_selected, ess_values = [], [], [] + modes = {key: Counter() for key in ("all_K", "selected_B", "Dplus", "Dminus")} + bounded_traces = [] + trap_active = defaultdict(bool) + policy_hash = BX.policy_sha256(policy) + + for step in range(int(cfg.T)): + start = time.perf_counter() + live = [replica for replica in replicas if replica.alive] + live, batch = BX._stack_prepared(live, device) + timers["sfm_stepping"] += time.perf_counter() - start + if not live: + break + counts["contexts"] += len(live) + + start = time.perf_counter() + with torch.no_grad(): + windows, contexts, x0 = _keyed_windows( + policy, live, batch, K=cfg.K, round_i=round_i, step=step, + source="K", seed=cfg.seed, nfe=cfg.nfe, temp=cfg.temp, + ) + raw_windows, _, raw_x0 = _keyed_windows( + policy, live, batch, K=1, round_i=round_i, step=step, + source="raw_continuation", seed=cfg.seed, + nfe=cfg.nfe, temp=cfg.temp, + ) + raw_windows = raw_windows[:, 0] + raw_x0 = raw_x0[:, 0] + windows_np = windows.detach().cpu().numpy() + raw_windows_np = raw_windows.detach().cpu().numpy() + x0_np = x0.detach().cpu().numpy() + raw_x0_np = raw_x0.detach().cpu().numpy() + timers["flow_proposal"] += time.perf_counter() - start + + start = time.perf_counter() + with torch.no_grad(): + features = _features_from_x0( + phi_policy, windows, contexts, x0, cfg.phi_s, + ) + raw_features = BR.l2_normalize(phi_policy.phi_s_from_x0( + raw_windows, contexts, raw_x0, s=cfg.phi_s, + )) + selected_by_context = [] + acquisitions = [] + for context_index, replica in enumerate(live): + generator = torch.Generator(device=features.device).manual_seed( + _keyed_seed( + cfg.seed, round_i, replica.scenario_id, + f"{replica.gamma:.8f}", step, "acquisition", + ) + ) + selected, trace = gp.sequential_acquire( + features[context_index], cfg.B, beta, generator=generator, + ) + selected_by_context.append(selected) + acquisitions.append(trace) + sigma_all.extend( + float(value) + for value in trace[0]["scores"].clamp_min(0.0).sqrt() + ) + sigma_selected.extend(float(row["chosen_sigma"]) for row in trace) + ess_values.extend(float(row["ess_norm"]) for row in trace) + timers["phi_rbf"] += time.perf_counter() - start + + tasks = [] + for context_index, replica in enumerate(live): + prepared = replica.prepared + for candidate_id in selected_by_context[context_index]: + tasks.append(( + context_index, + candidate_id, + prepared["state"], + windows_np[context_index, candidate_id], + prepared["ped_xy"], + prepared["ped_vel"], + replica.gamma, + )) + start = time.perf_counter() + results = list(executor.map(SM.verify_in_worker, tasks)) + timers["verifier"] += time.perf_counter() - start + counts["B_queries"] += len(tasks) + by_context = defaultdict(dict) + for context_index, candidate_id, result in results: + by_context[int(context_index)][int(candidate_id)] = result + + prepared_rows = [] + raw_tasks = [] + for context_index, replica in enumerate(live): + prepared = replica.prepared + prediction = SM.predict_pedestrians( + prepared["ped_xy"], prepared["ped_vel"], cfg.H, + ) + all_rows = [] + for candidate_id in range(cfg.K): + controls = windows_np[context_index, candidate_id] + segment = SM.rollout_positions(prepared["state"], controls) + mode = BE.classify_candidate(segment, prediction) + modes["all_K"][mode] += 1 + all_rows.append(dict( + candidate_id=int(candidate_id), + controls=controls, + mode=mode, + )) + query_rows = [] + for acquisition_step, candidate_id in enumerate( + selected_by_context[context_index] + ): + result = by_context[context_index][candidate_id] + label = FA._result_label(result) + counts[f"B_{label}"] += 1 + mode = all_rows[candidate_id]["mode"] + modes["selected_B"][mode] += 1 + if result.get("resolved"): + query_rows.append(dict( + candidate_id=int(candidate_id), + acquisition_step=int(acquisition_step), + controls=all_rows[candidate_id]["controls"], + result=result, + mode=mode, + sigma=float( + acquisitions[context_index][ + acquisition_step + ]["chosen_sigma"] + ), + )) + chosen = BC.select_admissible( + query_rows, + selector=cfg.selector, + state=prepared["state"], + ped_xy=prepared["ped_xy"], + ped_vel=prepared["ped_vel"], + gamma=replica.gamma, + ) + prepared_rows.append((all_rows, query_rows, chosen)) + if chosen is None: + raw_tasks.append(( + context_index, + -1, + prepared["state"], + raw_windows_np[context_index], + prepared["ped_xy"], + prepared["ped_vel"], + replica.gamma, + )) + + start = time.perf_counter() + raw_results = list(executor.map(SM.verify_in_worker, raw_tasks)) + timers["verifier"] += time.perf_counter() - start + counts["raw_continuation_queries"] += len(raw_tasks) + for context_index, candidate_id, result in raw_results: + if not result.get("resolved"): + raise RuntimeError( + "executed raw continuation verifier failed; " + "aborting instead of executing or omitting an unlabeled context: " + f"{result.get('error')}" + ) + by_context[int(context_index)][int(candidate_id)] = result + + start = time.perf_counter() + for context_index, replica in enumerate(live): + prepared = replica.prepared + all_rows, query_rows, chosen = prepared_rows[context_index] + context_id = shard.add_context( + scenario_id=replica.scenario_id, + gamma=replica.gamma, + step=step, + state=prepared["state"], + hp10=prepared["hp10"].numpy(), + low5=prepared["low"].numpy(), + hist=prepared["hist"].numpy(), + ped_xy=prepared["ped_xy"], + ped_vel=prepared["ped_vel"], + ) + nvp_context = chosen is None + if nvp_context: + controls = raw_windows_np[context_index] + selected_x0 = raw_x0_np[context_index] + result = by_context[context_index][-1] + margin, _, _ = BC.nominal_hp_margin( + prepared["state"], controls[0], prepared["ped_xy"], + replica.gamma, + ) + raw_admissible = bool( + result.get("resolved") + and int(result.get("y", 0)) == 1 + and bool(result.get("full_h")) + and float(margin) >= -1.0e-9 + ) + execution_source = ( + "certified_raw_rescue" + if raw_admissible else "uncertified_raw_continuation" + ) + candidate_id = None + acquisition_step = None + sigma = float(gp.acquisition_sigma( + raw_features[context_index:context_index + 1], + )[0]) + prediction = SM.predict_pedestrians( + prepared["ped_xy"], prepared["ped_vel"], cfg.H, + ) + mode = BE.classify_candidate( + SM.rollout_positions(prepared["state"], controls), + prediction, + ) + counts["NVP_contexts"] += 1 + else: + controls = np.asarray(chosen["controls"], np.float32) + selected_x0 = x0_np[context_index, int(chosen["candidate_id"])] + result = chosen["result"] + margin = float(chosen["hp_margin"]) + execution_source = f"verified_{cfg.selector}" + candidate_id = int(chosen["candidate_id"]) + acquisition_step = int(chosen["acquisition_step"]) + sigma = float(chosen["sigma"]) + mode = chosen["mode"] + + label = FA._result_label(result) + if label == "verifier_error": + raise RuntimeError( + "resolved executed-window partition requires an exact " + "full-H10 binary verifier label" + ) + counts[f"executed_{label}"] += 1 + counts[f"source_{execution_source}"] += 1 + window_id = shard.add_executed_window( + context_id, + controls, + selected_x0, + result, + execution_source=execution_source, + nvp_context=nvp_context, + candidate_id=candidate_id, + acquisition_step=acquisition_step, + sigma=sigma, + hp_margin=margin, + mode=mode, + ) + modes["Dplus" if int(result["y"]) == 1 else "Dminus"][mode] += 1 + + BX._advance(replica, controls[0]) + trap_event = FA._trap(replica.states) + trap_key = (replica.scenario_id, replica.gamma) + trap_entry = bool(trap_event and not trap_active[trap_key]) + trap_active[trap_key] = bool(trap_event) + collision, success, _ = FA._post_action_terminal(replica) + counts["trap_steps"] += int(trap_event) + counts["trap_entries"] += int(trap_entry) + counts["collision_events"] += int(collision) + counts["success_events"] += int(success) + if window_id is not None: + stored = shard.windows[int(window_id)] + stored.update( + trap_event=bool(trap_event), + trap_entry=bool(trap_entry), + collision_after_action=bool(collision), + success_after_action=bool(success), + ) + if len(bounded_traces) < 64 or ( + nvp_context and len(bounded_traces) < 128 + ): + bounded_traces.append(dict( + step=int(step), + scenario_id=int(replica.scenario_id), + gamma=float(replica.gamma), + selected_ids=list(map( + int, selected_by_context[context_index], + )), + B_labels=[ + FA._result_label( + by_context[context_index][candidate] + ) + for candidate in selected_by_context[context_index] + ], + execution_source=execution_source, + executed_label=label, + nvp_context=bool(nvp_context), + window_id=window_id, + trap_event=bool(trap_event), + collision_after_action=bool(collision), + success_after_action=bool(success), + )) + timers["sfm_stepping"] += time.perf_counter() - start + + _finalize_alive(replicas) + if BX.policy_sha256(policy) != policy_hash: + raise RuntimeError("policy changed during frozen offline macro-round") + shard_summary = shard.validate() + if ( + int(shard_summary["D"]) != int(shard_summary["contexts"]) + or int(shard_summary["errors"]) != 0 + or int(shard_summary["unresolved_contexts"]) != 0 + ): + raise RuntimeError( + "completed offline round must exactly partition every context " + "into D+ or D-" + ) + if int(counts["B_queries"]) != int(counts["contexts"]) * int(cfg.B): + raise RuntimeError("B verifier query accounting mismatch") + return dict( + collector_role="offline_expansion_data_collector_not_safe_controller", + continuation_semantics=( + f"verified {cfg.selector} B action when available; otherwise an " + "independent raw temperature-one H10 plan is exact-verified and " + "its first action is executed even when y=0" + ), + label_semantics=( + "D contains one resolved executed proposal per context; " + "D+=exact full-H10 y=1; D-=exact full-H10 y=0; " + "NVP/trap/collision are metadata and never retroactively relabel y" + ), + timers=dict(timers), + counts=dict(counts), + shard=shard_summary, + beta=float(beta), + realized_normalized_ess_over_remaining=float(np.mean(ess_values)), + sigma=BR.acquisition_diagnostics(sigma_all, sigma_selected), + modes={key: dict(value) for key, value in modes.items()}, + trace_examples=bounded_traces, + outcomes=[dict( + scenario_id=replica.scenario_id, + gamma=replica.gamma, + status=replica.status, + success=replica.status == "success", + collision=replica.status == "collision", + timeout=replica.status == "timeout", + steps=len(replica.controls), + min_clearance=float(replica.minimum_clearance), + ) for replica in replicas], + ) + + +def run(checkpoint, outdir, cfg, *, device): + cfg.validate() + checkpoint = os.path.abspath(checkpoint) + outdir = os.path.abspath(outdir) + if not os.path.isfile(checkpoint): + raise FileNotFoundError(checkpoint) + checkpoint_sha = OS.sha256_file(checkpoint) + if checkpoint_sha != EXPECTED_CHECKPOINT_SHA256: + raise ValueError( + f"checkpoint SHA mismatch: expected {EXPECTED_CHECKPOINT_SHA256}, " + f"got {checkpoint_sha}" + ) + if os.path.exists(outdir): + raise FileExistsError(f"refusing to reuse output directory: {outdir}") + os.makedirs(outdir) + environment = SS.scene_profile(cfg.scene_profile) + policy, _ = GPS.load_sfm_policy(checkpoint, device=device) + frozen_parameters = BS.configure_expansion_trainability(policy) + visual_encoder_sha = BS.module_sha256(policy.enc_grid) + optimizer = torch.optim.Adam( + [ + parameter for parameter in policy.parameters() + if parameter.requires_grad + ], + lr=cfg.lr, + ) + BX._save_checkpoint(policy, os.path.join(outdir, "round_00.pt"), dict( + round=0, + experiment=cfg.arm_name, + source_checkpoint=checkpoint, + source_sha256=checkpoint_sha, + encoder_sha256=visual_encoder_sha, + recipe=asdict(cfg), + )) + history = [] + previous_shard = None + preflight_scenarios = SP.expansion_scenarios(1, smoke=cfg.smoke) + preflight_replicas = [ + BX.Replica( + scenario_id, + gamma, + n_ped=environment["n_ped"], + ped_speed_range=tuple(environment["ped_speed_range"]), + ) + for scenario_id in preflight_scenarios for gamma in SP.GAMMAS + ] + ell0, ell, ell_preflight = _initial_lengthscale( + policy, preflight_replicas, cfg, device, + ) + with ProcessPoolExecutor(max_workers=cfg.verifier_workers) as executor: + for round_i in range(1, cfg.rounds + 1): + round_start = time.perf_counter() + scenarios = SP.expansion_scenarios(round_i, smoke=cfg.smoke) + replicas = [ + BX.Replica( + scenario_id, + gamma, + n_ped=environment["n_ped"], + ped_speed_range=tuple(environment["ped_speed_range"]), + ) + for scenario_id in scenarios for gamma in SP.GAMMAS + ] + if len(replicas) != 56: + raise RuntimeError("offline macro-round requires 56 episodes") + policy.eval() + phi_policy = copy.deepcopy(policy).eval() + for parameter in phi_policy.parameters(): + parameter.requires_grad_(False) + gp, gp_ids, gp_selection = gp_from_previous( + phi_policy, + previous_shard, + round_i=round_i, + ell=ell, + cap=CAP, + lam=cfg.gp_lam, + phi_s=cfg.phi_s, + device=device, + seed=cfg.seed + round_i * 101, + ) + beta, calibrated_ess = _calibrate_beta( + phi_policy, gp, replicas, cfg, device, round_i=round_i, + ) + shard = OS.ExecutedRoundShard(round_i) + gather = gather_offline_round( + policy, + phi_policy, + gp, + beta, + replicas, + cfg, + shard, + device, + executor, + round_i=round_i, + ) + shard_path = os.path.join( + outdir, "round_shards", f"round_{round_i:02d}.pt", + ) + shard_manifest = shard.save(shard_path) + replay_start = time.perf_counter() + replay = OR.replay( + policy, + optimizer, + shard, + alpha=cfg.alpha, + exposure_epochs=cfg.exposure_epochs, + batch=cfg.batch, + device=device, + seed=cfg.seed + round_i * 1_000_003, + ) + gather["timers"]["replay"] = time.perf_counter() - replay_start + if BS.module_sha256(policy.enc_grid) != visual_encoder_sha: + raise RuntimeError("visual encoder SHA changed") + checkpoint_path = os.path.join( + outdir, f"round_{round_i:02d}.pt", + ) + BX._save_checkpoint(policy, checkpoint_path, dict( + round=round_i, + experiment=cfg.arm_name, + source_checkpoint=checkpoint, + source_sha256=checkpoint_sha, + encoder_sha256=visual_encoder_sha, + recipe=asdict(cfg), + ell=ell, + ell0=ell0, + cap=CAP, + beta=float(beta), + )) + record = dict( + round=round_i, + experiment=cfg.arm_name, + scenarios=list(map(int, scenarios)), + environment=environment, + beta=float(beta), + calibrated_normalized_ess_over_remaining=float(calibrated_ess), + verifier=SM.verifier_manifest(), + gp_buffer_ids=gp_ids, + gp_selection=gp_selection, + gp=gp.diagnostics(), + gather=gather, + replay=replay, + shard=shard_manifest, + checkpoint=os.path.abspath(checkpoint_path), + checkpoint_sha256=OS.sha256_file(checkpoint_path), + wall_seconds=time.perf_counter() - round_start, + ) + history.append(record) + with open(os.path.join(outdir, "metrics.jsonl"), "a") as stream: + stream.write(json.dumps(record, allow_nan=False) + "\n") + print(json.dumps(dict( + round=round_i, + experiment=cfg.arm_name, + D=shard_manifest["D"], + Dplus=shard_manifest["Dplus"], + Dminus=shard_manifest["Dminus"], + beta=float(beta), + ess_over_remaining=float( + gather["realized_normalized_ess_over_remaining"] + ), + Adam_steps=int(replay["optimizer_steps"]), + wall_seconds=record["wall_seconds"], + )), flush=True) + previous_shard = shard + + manifest = dict( + status="SFM_B1_OFFLINE_EXEC_COMPLETE", + experiment=cfg.arm_name, + scientific_role="offline_expansion_data_collector_not_safe_controller", + recipe=asdict(cfg), + constants=dict( + ell=ell, + ell0=ell0, + ell_preflight=ell_preflight, + gp_buffer_cap=CAP, + gp_lambda=GP_LAMBDA, + expected_checkpoint_sha256=EXPECTED_CHECKPOINT_SHA256, + replay_window_rounds=1, + gp_quota_semantics=( + "exactly 73 executed D+ rows per gamma plus one rotating " + "extra; any support shortage aborts the scientific round" + ), + ess_target_semantics=( + "mean normalized ESS over each sequential remaining pool" + ), + ), + source=_source(), + source_checkpoint=checkpoint, + source_checkpoint_sha256=checkpoint_sha, + environment=environment, + frozen_parameters=frozen_parameters, + visual_encoder_sha=visual_encoder_sha, + history=history, + ) + _write_json(os.path.join(outdir, "COMPLETE.json"), manifest) + return manifest + + +def main(argv=None): + parser = argparse.ArgumentParser() + parser.add_argument("--checkpoint", required=True) + parser.add_argument("--outdir", required=True) + parser.add_argument("--alpha", type=float, choices=ALPHAS, required=True) + parser.add_argument( + "--selector", choices=EXECUTION_SELECTORS, default="margin", + ) + parser.add_argument( + "--exposure-epochs", + type=int, + choices=EXPOSURE_EPOCHS, + required=True, + ) + parser.add_argument("--rounds", type=int, default=10) + parser.add_argument("--verifier-workers", type=int, default=8) + parser.add_argument("--seed", type=int, default=20260724) + parser.add_argument("--device", default="cuda") + parser.add_argument("--smoke", action="store_true") + args = parser.parse_args(argv) + cfg = OfflineConfig( + alpha=args.alpha, + exposure_epochs=args.exposure_epochs, + selector=args.selector, + rounds=args.rounds, + verifier_workers=args.verifier_workers, + seed=args.seed, + smoke=args.smoke, + ) + run(args.checkpoint, args.outdir, cfg, device=args.device) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/sfm_b1_offline_replay.py b/overnight_run_07_12_sfm/sfm_b1_offline_replay.py new file mode 100644 index 0000000..37ee82b --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_offline_replay.py @@ -0,0 +1,326 @@ +"""Minibatch replay for one offline executed-window round. + +One exposure epoch visits every resolved executed positive and negative exactly +once. Adam steps after each mixed minibatch, rather than after accumulating a +single full-dataset gradient. +""" +from __future__ import annotations + +import math +import random + +import numpy as np +import torch + +import sfm_b1_r2_alpha_replay as R2 +import sfm_b1_store as BS +import sfm_b1_offline_store as OS + + +def _set_seed(seed): + seed = int(seed) + random.seed(seed) + np.random.seed(seed % (2 ** 32)) + torch.manual_seed(seed) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(seed) + + +def _identity(record): + shard, window = record + return int(shard.round_i), int(window["window_id"]) + + +def _proportional_interleave(positive, negative): + """Spread both deterministic sign orders over the complete epoch.""" + positive = list(positive) + negative = list(negative) + p_index = n_index = 0 + merged = [] + while p_index < len(positive) or n_index < len(negative): + p_progress = ( + (p_index + 1) / len(positive) + if p_index < len(positive) else float("inf") + ) + n_progress = ( + (n_index + 1) / len(negative) + if n_index < len(negative) else float("inf") + ) + if p_progress <= n_progress: + merged.append(positive[p_index]) + p_index += 1 + else: + merged.append(negative[n_index]) + n_index += 1 + return merged + + +def stratified_batches(shard, *, batch, seed): + batch = int(batch) + if batch < 1: + raise ValueError("batch must be positive") + positives = BS.hierarchical_order(OS.positive_records(shard), int(seed)) + negatives = BS.hierarchical_order(OS.negative_records(shard), int(seed) + 1) + total = len(positives) + len(negatives) + batch_count = math.ceil(total / batch) if total else 0 + if positives and len(positives) < batch_count: + raise RuntimeError( + "cannot place a positive in every fixed-capacity minibatch: " + f"{len(positives)} positives for {batch_count} batches" + ) + if not positives: + batches = [ + negatives[start:start + batch] + for start in range(0, len(negatives), batch) + ] + else: + # Signed replay needs a positive objective in every Adam step. Seed + # every fixed-capacity batch with one positive, then distribute the + # remaining deterministic sign orders without duplicating support. + batches = [[positives[index]] for index in range(batch_count)] + remaining = _proportional_interleave( + positives[batch_count:], negatives, + ) + batch_index = 0 + for record in remaining: + while len(batches[batch_index]) >= batch: + batch_index = (batch_index + 1) % batch_count + batches[batch_index].append(record) + batch_index = (batch_index + 1) % batch_count + identities = [_identity(record) for values in batches for record in values] + expected = [_identity(record) for record in positives + negatives] + if len(identities) != len(set(identities)) or set(identities) != set(expected): + raise RuntimeError("offline minibatch replay duplicated or omitted support") + if positives and any( + not any(int(record[1]["y"]) == 1 for record in values) + for values in batches + ): + raise RuntimeError("offline replay produced a positive-free minibatch") + if any(len(values) > batch for values in batches): + raise RuntimeError("offline replay exceeded the fixed minibatch capacity") + return batches, positives, negatives + + +def _weighted_loss(policy, records, mass, population, device): + if not records: + return None + grid, low, hist, controls = BS._tensor_batch(records, device) + context = policy.ctx_from(grid, low, hist) + weights = torch.as_tensor([ + int(population) * mass[(id(shard), int(window["query_id"]))] + for shard, window in records + ], dtype=controls.dtype, device=controls.device) + return policy.cfm_loss(controls, context, weights=weights) + + +def _finite_trainable(policy): + return all( + bool(torch.isfinite(parameter).all()) + for parameter in policy.parameters() + if parameter.requires_grad + ) + + +def _one_batch( + policy, optimizer, values, *, positive_mass, negative_mass, + positive_population, negative_population, alpha, device, seed, +): + positive = [record for record in values if int(record[1]["y"]) == 1] + negative = [record for record in values if int(record[1]["y"]) == 0] + _set_seed(seed) + optimizer.zero_grad(set_to_none=True) + positive_loss = _weighted_loss( + policy, positive, positive_mass, positive_population, device, + ) + if positive_loss is None: + return dict( + stepped=False, positive_loss=None, negative_loss=None, rho=0.0, + positive_norm=0.0, negative_norm=0.0, gradient_cosine=None, + positive=len(positive), negative=len(negative), + ) + if not bool(torch.isfinite(positive_loss)): + raise FloatingPointError("non-finite positive CFM loss") + positive_loss.backward() + positive_gradient = BS._gradient_snapshot(policy) + positive_norm = BS._gradient_norm(positive_gradient) + + negative_loss = None + negative_gradient = {} + negative_norm = 0.0 + rho = 0.0 + cosine = None + if float(alpha) > 0.0 and negative: + optimizer.zero_grad(set_to_none=True) + negative_loss = _weighted_loss( + policy, negative, negative_mass, negative_population, device, + ) + if not bool(torch.isfinite(negative_loss)): + raise FloatingPointError("non-finite negative CFM loss") + negative_loss.backward() + negative_gradient = BS._gradient_snapshot(policy) + negative_norm = BS._gradient_norm(negative_gradient) + rho = float(alpha) * positive_norm / (negative_norm + 1.0e-12) + cosine = R2._gradient_cosine(positive_gradient, negative_gradient) + + for name, parameter in policy.named_parameters(): + if not parameter.requires_grad: + continue + pos = positive_gradient.get(name) + neg = negative_gradient.get(name) + if pos is None and neg is None: + parameter.grad = None + elif pos is None: + parameter.grad = -rho * neg + elif neg is None: + parameter.grad = pos + else: + parameter.grad = pos - rho * neg + if parameter.grad is not None and not bool(torch.isfinite(parameter.grad).all()): + raise FloatingPointError(f"non-finite gradient in {name}") + optimizer.step() + if not _finite_trainable(policy): + raise FloatingPointError("optimizer produced non-finite parameters") + return dict( + stepped=True, + positive_loss=float(positive_loss.detach()), + negative_loss=( + None if negative_loss is None else float(negative_loss.detach()) + ), + rho=float(rho), + positive_norm=float(positive_norm), + negative_norm=float(negative_norm), + gradient_cosine=cosine, + positive=len(positive), + negative=len(negative), + ) + + +def _summary(values): + finite = [float(value) for value in values if value is not None] + if not finite: + return None + return dict( + mean=float(np.mean(finite)), + first=finite[0], + last=finite[-1], + minimum=min(finite), + maximum=max(finite), + ) + + +def replay( + policy, optimizer, shard, *, alpha, exposure_epochs, batch, device, seed, +): + if float(alpha) not in (0.0, 0.01, 0.1): + raise ValueError("alpha must be one of {0,0.01,0.1}") + if int(exposure_epochs) not in (1, 10, 100): + raise ValueError("exposure_epochs must be one of {1,10,100}") + policy.train() + positives = OS.positive_records(shard) + negatives = OS.negative_records(shard) + positive_mass, positive_mass_accounting = BS.hierarchy_mass(positives) + negative_mass, negative_mass_accounting = BS.hierarchy_mass(negatives) + probe_seed = int(seed) + 9_000_001 + fixed_probe_before = dict( + positive=R2._fixed_probe_loss( + policy, positives, batch=batch, device=device, seed=probe_seed, + ), + negative=R2._fixed_probe_loss( + policy, negatives, batch=batch, device=device, seed=probe_seed + 1, + ), + ) + module_before = R2._module_snapshot(policy) + encoder_before = BS.module_sha256(policy.enc_grid) + epoch_rows = [] + total_steps = 0 + for epoch_i in range(int(exposure_epochs)): + epoch_seed = int(seed) + epoch_i * 100_003 + batches, positive_order, negative_order = stratified_batches( + shard, batch=batch, seed=epoch_seed, + ) + batch_rows = [] + for batch_i, values in enumerate(batches): + row = _one_batch( + policy, optimizer, values, + positive_mass=positive_mass, + negative_mass=negative_mass, + positive_population=len(positives), + negative_population=len(negatives), + alpha=alpha, + device=device, + seed=epoch_seed + batch_i, + ) + batch_rows.append(row) + total_steps += int(row["stepped"]) + positive_visits = sum(row["positive"] for row in batch_rows) + negative_visits = sum(row["negative"] for row in batch_rows) + if positive_visits != len(positive_order): + raise RuntimeError("positive exposure count mismatch") + if negative_visits != len(negative_order): + raise RuntimeError("negative exposure count mismatch") + epoch_rows.append(dict( + epoch=epoch_i + 1, + seed=epoch_seed, + batches=len(batches), + optimizer_steps=sum(int(row["stepped"]) for row in batch_rows), + positive_visits=positive_visits, + negative_visits=negative_visits, + positive_loss=_summary(row["positive_loss"] for row in batch_rows), + negative_loss=_summary(row["negative_loss"] for row in batch_rows), + rho=_summary(row["rho"] for row in batch_rows), + positive_norm=_summary(row["positive_norm"] for row in batch_rows), + negative_norm=_summary(row["negative_norm"] for row in batch_rows), + gradient_cosine=_summary( + row["gradient_cosine"] for row in batch_rows + ), + )) + + encoder_after = BS.module_sha256(policy.enc_grid) + if encoder_after != encoder_before: + raise RuntimeError("visual encoder changed during offline replay") + fixed_probe_after = dict( + positive=R2._fixed_probe_loss( + policy, positives, batch=batch, device=device, seed=probe_seed, + ), + negative=R2._fixed_probe_loss( + policy, negatives, batch=batch, device=device, seed=probe_seed + 1, + ), + ) + expected_batches = ( + math.ceil((len(positives) + len(negatives)) / int(batch)) + if positives or negatives else 0 + ) + if any(row["batches"] != expected_batches for row in epoch_rows): + raise RuntimeError("unexpected minibatch count") + expected_steps = expected_batches * int(exposure_epochs) if positives else 0 + if total_steps != expected_steps: + raise RuntimeError( + f"expected {expected_steps} Adam steps, observed {total_steps}" + ) + policy.eval() + return dict( + alpha=float(alpha), + exposure_epochs=int(exposure_epochs), + positive_eligible=len(positives), + negative_eligible=len(negatives), + total_eligible=len(positives) + len(negatives), + batches_per_epoch=expected_batches, + optimizer_steps=total_steps, + positive_total_visits=len(positives) * int(exposure_epochs), + negative_total_visits=len(negatives) * int(exposure_epochs), + negative_used_for_training=bool(float(alpha) > 0.0 and negatives), + alpha_zero_semantics=( + "D- remains stored and occupies its deterministic mixed-minibatch " + "slots, but contributes exactly zero gradient when alpha=0" + ), + fixed_probe=dict(before=fixed_probe_before, after=fixed_probe_after), + module_relative_parameter_drift=R2._module_relative_drift( + module_before, R2._module_snapshot(policy), + ), + visual_encoder_sha_before=encoder_before, + visual_encoder_sha_after=encoder_after, + positive_mass=R2._compact_mass(positive_mass_accounting), + negative_mass=R2._compact_mass(negative_mass_accounting), + exact_once_per_exposure_epoch=True, + epochs=epoch_rows, + ) diff --git a/overnight_run_07_12_sfm/sfm_b1_offline_store.py b/overnight_run_07_12_sfm/sfm_b1_offline_store.py new file mode 100644 index 0000000..de90c2f --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_offline_store.py @@ -0,0 +1,230 @@ +"""Round-sharded store for offline executed-window SFM expansion. + +Each context contributes at most one resolved H=10 plan: the plan whose first +action was physically executed. The unexecuted B-1 verifier queries are search +diagnostics and never enter this store. +""" +from __future__ import annotations + +import hashlib +import json +import os + +import numpy as np +import torch + + +class ExecutedRoundShard: + VERSION = 2 + + def __init__(self, round_i): + self.round_i = int(round_i) + self.contexts = [] + self.windows = [] + self.errors = [] + self._context_keys = set() + self._window_contexts = set() + + def add_context( + self, *, scenario_id, gamma, step, state, hp10, low5, hist, + ped_xy, ped_vel, + ): + key = (int(scenario_id), round(float(gamma), 8), int(step)) + if key in self._context_keys: + raise ValueError(f"context already stored: {key}") + self._context_keys.add(key) + context_id = len(self.contexts) + self.contexts.append(dict( + context_id=context_id, + round=self.round_i, + scenario_id=int(scenario_id), + gamma=float(gamma), + step=int(step), + state=np.asarray(state, np.float32), + hp10=np.asarray(hp10, np.float32), + low5=np.asarray(low5, np.float32), + hist=np.asarray(hist, np.float32), + ped_xy=np.asarray(ped_xy, np.float32), + ped_vel=np.asarray(ped_vel, np.float32), + )) + return context_id + + def add_executed_window( + self, context_id, controls, x0, result, *, execution_source, nvp_context, + candidate_id=None, acquisition_step=None, sigma=None, hp_margin=None, + mode=None, + ): + if not bool(result.get("resolved")): + raise ValueError("unresolved executed verifier result cannot enter D") + if int(result.get("y", -1)) not in (0, 1): + raise ValueError("resolved executed window needs binary y") + if not bool(result.get("full_h")) or int(result.get("terminal_step", -1)) != 10: + raise ValueError("offline training D requires exact full-H=10 labels") + controls = np.asarray(controls, np.float32) + if tuple(controls.shape) != (10, 2) or not np.isfinite(controls).all(): + raise ValueError("offline training D requires finite controls [10,2]") + x0 = np.asarray(x0, np.float32) + if tuple(x0.shape) != (20,) or not np.isfinite(x0).all(): + raise ValueError("offline training D requires finite original x0 [20]") + components = ( + bool(result.get("taskspace")), + bool(result.get("collision_free")), + bool(result.get("certificate")), + ) + if int(result["y"]) != int(all(components)): + raise ValueError("verifier y disagrees with its exact components") + context_id = int(context_id) + if not 0 <= context_id < len(self.contexts): + raise IndexError("unknown context") + if context_id in self._window_contexts: + raise ValueError("a context can contribute at most one executed window") + self._window_contexts.add(context_id) + window_id = len(self.windows) + self.windows.append(dict( + window_id=window_id, + query_id=window_id, + context_id=context_id, + controls=controls, + x0=x0, + y=int(result["y"]), + taskspace=bool(result["taskspace"]), + collision_free=bool(result["collision_free"]), + certificate=bool(result["certificate"]), + full_h=True, + terminal_step=10, + train_eligible=bool(result["y"]), + execution_source=str(execution_source), + nvp_context=bool(nvp_context), + candidate_id=None if candidate_id is None else int(candidate_id), + acquisition_step=( + None if acquisition_step is None else int(acquisition_step) + ), + sigma=None if sigma is None else float(sigma), + hp_margin=None if hp_margin is None else float(hp_margin), + mode=mode, + verifier_diagnostics=dict(result["diagnostics"]), + )) + return window_id + + def add_error(self, *, context_id, candidate_id, execution_source, error): + self.errors.append(dict( + context_id=int(context_id), + candidate_id=None if candidate_id is None else int(candidate_id), + execution_source=str(execution_source), + error=str(error), + )) + + @property + def D(self): + return list(self.windows) + + @property + def Dplus(self): + return [row for row in self.windows if row["y"] == 1] + + @property + def Dminus(self): + return [row for row in self.windows if row["y"] == 0] + + def validate(self): + for expected, context in enumerate(self.contexts): + if int(context["context_id"]) != expected: + raise AssertionError("context IDs are not dense") + seen = set() + for expected, window in enumerate(self.windows): + if ( + int(window["window_id"]) != expected + or int(window["query_id"]) != expected + ): + raise AssertionError("window IDs are not dense") + context_id = int(window["context_id"]) + if not 0 <= context_id < len(self.contexts): + raise AssertionError("window references missing context") + if context_id in seen: + raise AssertionError("multiple training windows share one context") + seen.add(context_id) + if not window["full_h"] or int(window["terminal_step"]) != 10: + raise AssertionError("non-H10 window entered training D") + x0 = np.asarray(window.get("x0"), np.float32) + if tuple(x0.shape) != (20,) or not np.isfinite(x0).all(): + raise AssertionError("window is missing its finite original x0 [20]") + if len(self.Dplus) + len(self.Dminus) != len(self.D): + raise AssertionError("D+/D- must exactly partition resolved executed D") + return dict( + round=self.round_i, + contexts=len(self.contexts), + D=len(self.D), + Dplus=len(self.Dplus), + Dminus=len(self.Dminus), + errors=len(self.errors), + unresolved_contexts=len(self.contexts) - len(self.D), + ) + + def save(self, path): + path = os.path.abspath(os.fspath(path)) + summary = self.validate() + os.makedirs(os.path.dirname(path), exist_ok=True) + temporary = path + ".tmp" + torch.save(dict( + version=self.VERSION, + round=self.round_i, + contexts=self.contexts, + windows=self.windows, + errors=self.errors, + summary=summary, + ), temporary) + os.replace(temporary, path) + digest = sha256_file(path) + complete = path + ".COMPLETE.json" + with open(complete + ".tmp", "w") as stream: + json.dump(dict( + status="OFFLINE_EXECUTED_ROUND_SHARD_COMPLETE", + file=path, + sha256=digest, + **summary, + ), stream, indent=2) + os.replace(complete + ".tmp", complete) + return dict(path=path, sha256=digest, complete=complete, **summary) + + @classmethod + def load(cls, path): + payload = torch.load(path, map_location="cpu", weights_only=False) + if int(payload["version"]) != cls.VERSION: + raise ValueError("unsupported executed-round shard version") + value = cls(payload["round"]) + value.contexts = payload["contexts"] + value.windows = payload["windows"] + value.errors = payload["errors"] + value._context_keys = { + ( + int(row["scenario_id"]), + round(float(row["gamma"]), 8), + int(row["step"]), + ) + for row in value.contexts + } + value._window_contexts = { + int(row["context_id"]) for row in value.windows + } + value.validate() + return value + + +def context_for(shard, window): + return shard.contexts[int(window["context_id"])] + + +def positive_records(shard): + return [(shard, row) for row in shard.Dplus] + + +def negative_records(shard): + return [(shard, row) for row in shard.Dminus] + + +def sha256_file(path): + digest = hashlib.sha256() + with open(path, "rb") as stream: + for chunk in iter(lambda: stream.read(1 << 20), b""): + digest.update(chunk) + return digest.hexdigest() diff --git a/overnight_run_07_12_sfm/sfm_b1_r2_aggregate.py b/overnight_run_07_12_sfm/sfm_b1_r2_aggregate.py new file mode 100644 index 0000000..8a07f20 --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_r2_aggregate.py @@ -0,0 +1,341 @@ +#!/usr/bin/env python3 +"""Aggregate the completed two-round SFM alpha/replay factorial. + +The per-arm evaluator remains the source of scientific metrics. This module +only validates their shared contracts, renders pooled comparisons, and +computes paired scenario-bootstrap changes against the common r0 checkpoint. +""" +from __future__ import annotations + +import argparse +import csv +import hashlib +import json +import math +import os +from pathlib import Path + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np + + +ALPHAS = (0.0, 0.01, 0.1) +REPLAY_EPOCHS = (1, 10, 100) +ROUNDS = (0, 1, 2) + + +def arm_name(alpha: float, epochs: int) -> str: + return f"margin_alpha{str(float(alpha)).replace('.', 'p')}_epochs{int(epochs):03d}" + + +def _sha256_file(path: str | os.PathLike[str]) -> str: + digest = hashlib.sha256() + with open(path, "rb") as stream: + for chunk in iter(lambda: stream.read(1 << 20), b""): + digest.update(chunk) + return digest.hexdigest() + + +def _write_json(path: Path, value: dict) -> None: + temporary = path.with_name(path.name + ".tmp") + with temporary.open("w") as stream: + json.dump(value, stream, indent=2, allow_nan=False) + os.replace(temporary, path) + + +def _metric(cell: dict, key: str) -> float | None: + if key in ("SR", "CR", "timeout", "V_safe"): + return float(cell[key]) + entry = ( + cell["successful_clearance"] + if key == "clearance" + else cell["successful_time_to_goal"] + ) + return None if entry["mean"] is None else float(entry["mean"]) + + +def _post_expansion_key(row: dict) -> tuple: + """Safety-first deterministic selection among measured r1/r2 cells.""" + clearance = -math.inf if row["clearance"] is None else float(row["clearance"]) + time = math.inf if row["time"] is None else float(row["time"]) + return ( + float(row["CR"]), + -float(row["SR"]), + -clearance, + time, + int(row["round"]), + float(row["alpha"]), + int(row["replay_epochs"]), + ) + + +def paired_cluster_delta( + baseline_rows: list[dict], + candidate_rows: list[dict], + key: str, + *, + seed: int = 20260723, + draws: int = 20_000, +) -> dict: + baseline = { + (int(row["episode"]), float(row["gamma"])): float(bool(row[key])) + for row in baseline_rows + } + candidate = { + (int(row["episode"]), float(row["gamma"])): float(bool(row[key])) + for row in candidate_rows + } + if set(baseline) != set(candidate): + raise ValueError("paired evaluator rows do not share the same scenario/gamma keys") + episode_ids = sorted({episode for episode, _ in baseline}) + per_episode = np.asarray([ + np.mean([ + candidate[(episode, gamma)] - baseline[(episode, gamma)] + for current_episode, gamma in baseline if current_episode == episode + ]) + for episode in episode_ids + ], dtype=float) + generator = np.random.default_rng(int(seed)) + indices = generator.integers( + 0, len(per_episode), size=(int(draws), len(per_episode)) + ) + samples = per_episode[indices].mean(axis=1) + return { + "estimate": float(per_episode.mean()), + "scenario_cluster_bootstrap95": list( + map(float, np.quantile(samples, (0.025, 0.975))) + ), + "paired_scenarios": len(episode_ids), + "paired_gamma_cells": len(baseline), + } + + +def _load(run_root: Path) -> tuple[list[dict], dict, str]: + rows = [] + reference_payload = None + noise_sha = None + baseline_cell_key = None + for alpha in ALPHAS: + for epochs in REPLAY_EPOCHS: + arm = arm_name(alpha, epochs) + training_path = run_root / "arms" / arm / "COMPLETE.json" + evaluation_path = ( + run_root / "evaluation" / arm / "raw_m50_r0_r2_metrics.json" + ) + with training_path.open() as stream: + training = json.load(stream) + with evaluation_path.open() as stream: + evaluation = json.load(stream) + if training.get("status") != "R2_ALPHA_REPLAY_COMPLETE": + raise RuntimeError(f"incomplete training arm: {training_path}") + if evaluation.get("status") != "SFM_B1_R2_RAW_M50_COMPLETE": + raise RuntimeError(f"incomplete evaluation arm: {evaluation_path}") + current_noise = evaluation["noise_bank"]["sha256"] + if noise_sha is None: + noise_sha = current_noise + reference_payload = evaluation["archived_M100_reference"] + elif current_noise != noise_sha: + raise RuntimeError("arm evaluations do not share one CRN bank") + records = evaluation["records"] + if [int(record["round"]) for record in records] != list(ROUNDS): + raise RuntimeError(f"{arm} does not contain r0/r1/r2") + if baseline_cell_key is None: + baseline_cell_key = records[0]["cell"]["cell_key"] + elif records[0]["cell"]["cell_key"] != baseline_cell_key: + raise RuntimeError("arm evaluations do not reuse the same r0 cell") + history = {int(item["round"]): item for item in training["history"]} + for record in records: + round_i = int(record["round"]) + cell = record["cell"]["summary"]["pooled"] + item = { + "arm": arm, + "alpha": float(alpha), + "replay_epochs": int(epochs), + "round": round_i, + "SR": _metric(cell, "SR"), + "CR": _metric(cell, "CR"), + "timeout": _metric(cell, "timeout"), + "V_safe": _metric(cell, "V_safe"), + "clearance": _metric(cell, "clearance"), + "time": _metric(cell, "time"), + "cell_key": record["cell"]["cell_key"], + "checkpoint_sha256": record["cell"]["checkpoint_sha256"], + "evaluation_rows": record["cell"]["rows"], + "training": None if round_i == 0 else history[round_i], + } + rows.append(item) + assert reference_payload is not None and noise_sha is not None + return rows, reference_payload, noise_sha + + +def _render(rows: list[dict], outdir: Path, best: dict) -> list[str]: + specs = ( + ("CR", "Collision rate"), + ("V_safe", r"$V_{\mathrm{safe}}$"), + ("clearance", "Successful min. clearance [m]"), + ("time", "Successful time-to-goal [s]"), + ) + colors = {1: "#0072B2", 10: "#E69F00", 100: "#CC79A7"} + linestyles = {0.0: "-", 0.01: "--", 0.1: ":"} + plt.rcParams.update({ + "font.family": "serif", + "mathtext.fontset": "cm", + "font.serif": ["cmr10", "Computer Modern Roman", "DejaVu Serif"], + "axes.unicode_minus": False, + "axes.formatter.use_mathtext": True, + }) + figure, axes = plt.subplots(2, 2, figsize=(14.5, 9)) + for axis, (metric, title) in zip(axes.flat, specs): + for alpha in ALPHAS: + for epochs in REPLAY_EPOCHS: + values = sorted( + [ + row for row in rows + if row["alpha"] == alpha + and row["replay_epochs"] == epochs + ], + key=lambda row: row["round"], + ) + axis.plot( + [row["round"] for row in values], + [ + np.nan if row[metric] is None else row[metric] + for row in values + ], + color=colors[epochs], + linestyle=linestyles[alpha], + marker="o", + lw=2, + alpha=0.85, + ) + axis.scatter( + [best["round"]], + [best[metric]], + marker="*", + s=210, + c="#009E73", + edgecolors="black", + zorder=8, + ) + axis.set( + title=title, + xlabel="expansion round", + xticks=ROUNDS, + ) + axis.grid(alpha=0.25) + if metric in ("CR", "V_safe"): + axis.set_ylim(-0.03, 1.03) + handles = [ + plt.Line2D([0], [0], color=colors[value], lw=2.5, label=f"{value} epochs") + for value in REPLAY_EPOCHS + ] + handles.extend([ + plt.Line2D( + [0], [0], color="black", linestyle=linestyles[value], lw=2, + label=rf"$\alpha={value:g}$", + ) + for value in ALPHAS + ]) + handles.append(plt.Line2D( + [0], [0], marker="*", color="none", markerfacecolor="#009E73", + markeredgecolor="black", markersize=14, label="best post-expansion cell", + )) + figure.legend( + handles=handles, loc="upper center", ncol=7, frameon=False, + bbox_to_anchor=(0.5, 0.995), + ) + figure.tight_layout(rect=(0.02, 0.02, 0.98, 0.92)) + outputs = [] + for suffix in ("png", "pdf"): + path = outdir / f"factorial_pooled_curves.{suffix}" + figure.savefig(path, dpi=300, bbox_inches="tight") + outputs.append(str(path)) + plt.close(figure) + return outputs + + +def run(run_root: str) -> dict: + root = Path(run_root).resolve() + rows, archived, noise_sha = _load(root) + baseline = next( + row for row in rows + if row["round"] == 0 + and row["alpha"] == 0.0 + and row["replay_epochs"] == 1 + ) + candidates = [row for row in rows if row["round"] > 0] + best = min(candidates, key=_post_expansion_key) + output_dir = root / "evaluation" / "aggregate" + output_dir.mkdir(parents=True, exist_ok=True) + csv_path = output_dir / "factorial_pooled_metrics.csv" + fields = ( + "arm", "alpha", "replay_epochs", "round", + "SR", "CR", "timeout", "V_safe", "clearance", "time", + ) + with csv_path.open("w", newline="") as stream: + writer = csv.DictWriter(stream, fieldnames=fields) + writer.writeheader() + for row in sorted( + rows, key=lambda item: ( + item["alpha"], item["replay_epochs"], item["round"] + ) + ): + writer.writerow({key: row[key] for key in fields}) + outputs = _render(rows, output_dir, best) + summary = { + "status": "SFM_B1_R2_FACTORIAL_AGGREGATE_COMPLETE", + "run_root": str(root), + "training_source_commit": "58ec896f87a5859149a39f5f7796560cd53da518", + "noise_bank_sha256": noise_sha, + "common_r0_cell_key": baseline["cell_key"], + "baseline": {key: baseline[key] for key in fields}, + "best_post_expansion": {key: best[key] for key in fields}, + "paired_changes_best_minus_r0": { + key: paired_cluster_delta( + baseline["evaluation_rows"], + best["evaluation_rows"], + key, + seed=20260723 + index, + ) + for index, key in enumerate(("success", "collision", "timeout", "v_safe")) + }, + "archived_M100_reference": archived, + "selection_rule": ( + "post-expansion cells only; lower raw CR, then higher raw SR, then " + "higher successful-only clearance, then lower successful-only time" + ), + "artifacts": { + "csv": str(csv_path), + "csv_sha256": _sha256_file(csv_path), + "figures": [ + {"path": path, "sha256": _sha256_file(path)} + for path in outputs + ], + "best_per_gamma_png": str( + root / "evaluation" / best["arm"] / "raw_m50_r0_r2_curves.png" + ), + }, + } + summary_path = output_dir / "factorial_summary.json" + _write_json(summary_path, summary) + return summary + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--run-root", required=True) + result = run(parser.parse_args().run_root) + print(json.dumps({ + "status": result["status"], + "baseline": result["baseline"], + "best_post_expansion": result["best_post_expansion"], + "paired_changes": result["paired_changes_best_minus_r0"], + "artifacts": result["artifacts"], + }, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/sfm_b1_r2_alpha_replay.py b/overnight_run_07_12_sfm/sfm_b1_r2_alpha_replay.py new file mode 100644 index 0000000..a960f6a --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_r2_alpha_replay.py @@ -0,0 +1,561 @@ +"""Isolated two-round max-margin B1 alpha/replay-epoch experiment. + +This module intentionally does not change the authenticated Arm-A runner. It +keeps the B1 gather path fixed and varies only + + alpha in {0, 0.01, 0.1} + complete W=2 replay epochs in {1, 10, 100}. + +One replay epoch visits every eligible record exactly once, accumulates the +hierarchically weighted objective over minibatches, and then takes one Adam +step. Consequently ``replay_epochs`` is also the number of optimizer steps +per macro-round whenever positive support is non-empty. +""" +from __future__ import annotations + +import argparse +import copy +from dataclasses import asdict, dataclass +import json +import os +import random +import time +from concurrent.futures import ProcessPoolExecutor + +import numpy as np +import torch + +import _paths # noqa: F401 +import grid_policy_sfm as GPS +import sfm_b1_expand as BX +import sfm_b1_store as BS +import sfm_protocol as SP +import sfm_scene as SS +import sfm_metrics2 as SM + + +EXPECTED_CHECKPOINT_SHA256 = "1b5179c935d3eeff8824967d707d64cc9bab273949ee1f0e4f190172bab1b215" +ELL = 0.24210826720721101 +CAP = 256 +GP_LAMBDA = 1.0e-2 +LEARNING_RATE = 1.0e-4 +ROUNDS = 2 +ALPHAS = (0.0, 0.01, 0.1) +REPLAY_EPOCHS = (1, 10, 100) + + +@dataclass(frozen=True) +class ExperimentConfig: + alpha: float + replay_epochs: int + rounds: int = ROUNDS + K: int = 16 + B: int = 4 + T: int = 180 + H: int = 10 + W: int = 2 + batch: int = 128 + lr: float = LEARNING_RATE + ess_target: float = 0.5 + nfe: int = 8 + temp: float = 1.0 + phi_s: float = 0.9 + gp_lam: float = GP_LAMBDA + selector: str = "margin" + verifier_workers: int = 32 + seed: int = 20260723 + scene_profile: str = "double_density_velocity_ood" + smoke: bool = False + + def validate(self): + if float(self.alpha) not in ALPHAS: + raise ValueError(f"alpha must be one of {ALPHAS}") + if int(self.replay_epochs) not in REPLAY_EPOCHS: + raise ValueError(f"replay_epochs must be one of {REPLAY_EPOCHS}") + expected = ( + ROUNDS, 16, 4, 180, 10, 2, 128, LEARNING_RATE, 0.5, 8, 1.0, + 0.9, GP_LAMBDA, "margin", "double_density_velocity_ood", False, + ) + actual = ( + self.rounds, self.K, self.B, self.T, self.H, self.W, self.batch, + self.lr, self.ess_target, self.nfe, self.temp, self.phi_s, + self.gp_lam, self.selector, self.scene_profile, self.smoke, + ) + if actual != expected: + raise ValueError("fixed two-round alpha/replay experiment contract changed") + if int(self.verifier_workers) < 1: + raise ValueError("verifier_workers must be positive") + return self + + @property + def arm_name(self): + alpha = str(float(self.alpha)).replace(".", "p") + return f"margin_alpha{alpha}_epochs{int(self.replay_epochs):03d}" + + +def _set_update_seed(seed): + """Seed every RNG used by CFM noise, dropout, and replay ordering.""" + seed = int(seed) + random.seed(seed) + np.random.seed(seed % (2 ** 32)) + torch.manual_seed(seed) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(seed) + + +def _records(recent, alpha): + positives = [ + (shard, query) + for shard, query in recent.positive_records() + if query["train_eligible"] + ] + # Preserve the established alpha=0 contract exactly: D- is not read, + # including for diagnostics. Its negative fixed-probe loss is therefore + # explicitly reported as null in the alpha=0 arm. + negatives = [] if float(alpha) == 0.0 else list(recent.negative_records()) + return positives, negatives + + +def _coverage(visited, eligible): + identities = list(visited) + return dict( + eligible=int(eligible), + visited=len(identities), + unique_visited=len(set(identities)), + exact_once=bool( + len(identities) == int(eligible) + and len(set(identities)) == int(eligible) + ), + ) + + +def _compact_mass(accounting): + """Keep exact mass checks without serializing every cell/context key.""" + gamma = dict(accounting.get("gamma", {})) + gamma_values = list(map(float, gamma.values())) + return dict( + total=float(accounting.get("total", 0.0)), + gamma=gamma, + gamma_spread=( + max(gamma_values) - min(gamma_values) if gamma_values else 0.0 + ), + cells=len(accounting.get("cells", {})), + contexts=len(accounting.get("contexts", {})), + ) + + +def _gradient_cosine(left, right): + dot = torch.zeros((), dtype=torch.float64) + left_sq = torch.zeros((), dtype=torch.float64) + right_sq = torch.zeros((), dtype=torch.float64) + for name in set(left) | set(right): + lvalue = left.get(name) + rvalue = right.get(name) + if lvalue is not None: + left_sq += lvalue.to(torch.float64).square().sum().cpu() + if rvalue is not None: + right_sq += rvalue.to(torch.float64).square().sum().cpu() + if lvalue is not None and rvalue is not None: + dot += (lvalue.to(torch.float64) * rvalue.to(torch.float64)).sum().cpu() + denominator = float((left_sq * right_sq).sqrt()) + return None if denominator <= 0.0 else float(dot) / denominator + + +def _group_gradient_norms(policy, snapshot): + parameter_names = {id(parameter): name for name, parameter in policy.named_parameters()} + result = {} + for group_name, module in policy.module_groups().items(): + squared = torch.zeros((), dtype=torch.float64) + for parameter in module.parameters(): + value = snapshot.get(parameter_names[id(parameter)]) + if value is not None: + squared += value.to(torch.float64).square().sum().cpu() + result[group_name] = float(squared.sqrt()) + return result + + +def _module_snapshot(policy): + return { + group: { + name: value.detach().cpu().clone() + for name, value in module.state_dict().items() + } + for group, module in policy.module_groups().items() + } + + +def _module_relative_drift(before, after, eps=1.0e-12): + result = {} + for group in before: + delta_sq = torch.zeros((), dtype=torch.float64) + base_sq = torch.zeros((), dtype=torch.float64) + for name, initial in before[group].items(): + final = after[group][name] + delta_sq += (final.to(torch.float64) - initial.to(torch.float64)).square().sum() + base_sq += initial.to(torch.float64).square().sum() + result[group] = float(delta_sq.sqrt() / base_sq.sqrt().clamp_min(float(eps))) + return result + + +@torch.no_grad() +def _fixed_probe_loss(policy, records, *, batch, device, seed): + """Evaluate deterministic, dropout-disabled CFM loss on complete support.""" + if not records: + return None + mass, _ = BS.hierarchy_mass(records) + was_training = policy.training + policy.eval() + generator = torch.Generator(device=device).manual_seed(int(seed)) + total = 0.0 + for values in BS._batches(records, int(batch)): + grid, low, hist, controls = BS._tensor_batch(values, device) + context = policy.ctx_from(grid, low, hist) + count = len(values) + x1 = (controls / policy.u_max).reshape(count, policy.d) + x0 = torch.randn( + x1.shape, dtype=x1.dtype, device=x1.device, generator=generator, + ) + tau = torch.rand( + count, dtype=x1.dtype, device=x1.device, generator=generator, + ).clamp(1.0e-4, 1.0) + x_tau = (1.0 - tau)[:, None] * x0 + tau[:, None] * x1 + target = x1 - x0 + prediction = policy.forward(x_tau, tau, context) + per = (prediction - target).square().mean(dim=1) + weights = torch.as_tensor( + [mass[(id(shard), int(query["query_id"]))] for shard, query in values], + dtype=per.dtype, device=per.device, + ) + total += float((per * weights).sum()) + policy.train(was_training) + return float(total) + + +def _positive_epoch(policy, optimizer, positives, *, batch, device, seed, path): + ordered = BS.hierarchical_order(positives, seed) + mass, accounting = BS.hierarchy_mass(ordered) + optimizer.zero_grad(set_to_none=True) + if not ordered: + return dict( + path=path, optimizer_steps=0, positive_loss=0.0, + positive_norm=0.0, positive_group_norms={}, + positive_coverage=_coverage([], 0), + positive_mass=_compact_mass(accounting), + negative_coverage=_coverage([], 0), + ) + loss, visited = BS._accumulate_objective(policy, ordered, mass, batch, device) + gradient = BS._gradient_snapshot(policy) + norm = BS._gradient_norm(gradient) + group_norms = _group_gradient_norms(policy, gradient) + coverage = _coverage(visited, len(ordered)) + if not coverage["exact_once"]: + raise RuntimeError("positive replay did not cover every eligible record exactly once") + optimizer.step() + return dict( + path=path, optimizer_steps=1, positive_loss=float(loss), + positive_norm=float(norm), positive_group_norms=group_norms, + positive_coverage=coverage, positive_mass=_compact_mass(accounting), + negative_coverage=_coverage([], 0), + ) + + +def _signed_epoch(policy, optimizer, positives, negatives, *, alpha, batch, device, seed): + if float(alpha) == 0.0: + return _positive_epoch( + policy, optimizer, positives, batch=batch, device=device, + seed=seed, path="positive_only", + ) + if not negatives: + return _positive_epoch( + policy, optimizer, positives, batch=batch, device=device, + seed=seed, path="positive_fallback_no_negative", + ) + + positive_order = BS.hierarchical_order(positives, seed) + negative_order = BS.hierarchical_order(negatives, seed + 1) + positive_mass, positive_accounting = BS.hierarchy_mass(positive_order) + negative_mass, negative_accounting = BS.hierarchy_mass(negative_order) + if not positive_order: + optimizer.zero_grad(set_to_none=True) + negative_loss, negative_visited = BS._accumulate_objective( + policy, negative_order, negative_mass, batch, device, + ) + negative_gradient = BS._gradient_snapshot(policy) + optimizer.zero_grad(set_to_none=True) + negative_coverage = _coverage(negative_visited, len(negative_order)) + if not negative_coverage["exact_once"]: + raise RuntimeError("negative-only replay coverage failure") + return dict( + path="signed_no_positive", optimizer_steps=0, alpha=float(alpha), + rho=0.0, gradient_cosine=None, positive_loss=0.0, + negative_loss=float(negative_loss), positive_norm=0.0, + negative_norm=BS._gradient_norm(negative_gradient), + positive_group_norms={}, + negative_group_norms=_group_gradient_norms(policy, negative_gradient), + positive_coverage=_coverage([], 0), + negative_coverage=negative_coverage, + positive_mass=_compact_mass(positive_accounting), + negative_mass=_compact_mass(negative_accounting), + ) + + optimizer.zero_grad(set_to_none=True) + positive_loss, positive_visited = BS._accumulate_objective( + policy, positive_order, positive_mass, batch, device, + ) + positive_gradient = BS._gradient_snapshot(policy) + positive_norm = BS._gradient_norm(positive_gradient) + optimizer.zero_grad(set_to_none=True) + negative_loss, negative_visited = BS._accumulate_objective( + policy, negative_order, negative_mass, batch, device, + ) + negative_gradient = BS._gradient_snapshot(policy) + negative_norm = BS._gradient_norm(negative_gradient) + rho = float(alpha) * positive_norm / (negative_norm + 1.0e-12) + for name, parameter in policy.named_parameters(): + if not parameter.requires_grad: + continue + positive = positive_gradient.get(name) + negative = negative_gradient.get(name) + if positive is None and negative is None: + parameter.grad = None + elif positive is None: + parameter.grad = -rho * negative + elif negative is None: + parameter.grad = positive + else: + parameter.grad = positive - rho * negative + + positive_coverage = _coverage(positive_visited, len(positive_order)) + negative_coverage = _coverage(negative_visited, len(negative_order)) + if not positive_coverage["exact_once"] or not negative_coverage["exact_once"]: + raise RuntimeError("signed replay did not cover complete support exactly once") + optimizer.step() + return dict( + path="signed", optimizer_steps=1, alpha=float(alpha), rho=float(rho), + gradient_cosine=_gradient_cosine(positive_gradient, negative_gradient), + positive_loss=float(positive_loss), negative_loss=float(negative_loss), + positive_norm=float(positive_norm), negative_norm=float(negative_norm), + positive_group_norms=_group_gradient_norms(policy, positive_gradient), + negative_group_norms=_group_gradient_norms(policy, negative_gradient), + positive_coverage=positive_coverage, negative_coverage=negative_coverage, + positive_mass=_compact_mass(positive_accounting), + negative_mass=_compact_mass(negative_accounting), + ) + + +def _numeric_summary(rows, key): + values = [float(row[key]) for row in rows if row.get(key) is not None] + if not values: + return None + return dict( + first=values[0], last=values[-1], mean=float(np.mean(values)), + minimum=min(values), maximum=max(values), + ) + + +def repeat_complete_replay(policy, optimizer, recent, cfg, *, device, round_i): + positives, negatives = _records(recent, cfg.alpha) + probe_seed = cfg.seed + int(round_i) * 1_000_003 + fixed_probe_before = dict( + positive=_fixed_probe_loss( + policy, positives, batch=cfg.batch, device=device, seed=probe_seed, + ), + negative=_fixed_probe_loss( + policy, negatives, batch=cfg.batch, device=device, seed=probe_seed + 1, + ), + ) + modules_before = _module_snapshot(policy) + encoder_before = BS.module_sha256(policy.enc_grid) + epoch_rows = [] + for epoch_i in range(int(cfg.replay_epochs)): + epoch_seed = cfg.seed + int(round_i) * 100_000 + epoch_i + _set_update_seed(epoch_seed) + row = _signed_epoch( + policy, optimizer, positives, negatives, alpha=cfg.alpha, + batch=cfg.batch, device=device, seed=epoch_seed, + ) + epoch_rows.append(dict(epoch=epoch_i + 1, seed=epoch_seed, **row)) + encoder_after = BS.module_sha256(policy.enc_grid) + if encoder_after != encoder_before: + raise RuntimeError("visual encoder changed during isolated replay") + fixed_probe_after = dict( + positive=_fixed_probe_loss( + policy, positives, batch=cfg.batch, device=device, seed=probe_seed, + ), + negative=_fixed_probe_loss( + policy, negatives, batch=cfg.batch, device=device, seed=probe_seed + 1, + ), + ) + modules_after = _module_snapshot(policy) + optimizer_steps = sum(int(row["optimizer_steps"]) for row in epoch_rows) + expected_steps = int(cfg.replay_epochs) if positives else 0 + if optimizer_steps != expected_steps: + raise RuntimeError( + f"expected {expected_steps} optimizer steps, observed {optimizer_steps}" + ) + if any(not row["positive_coverage"]["exact_once"] for row in epoch_rows if positives): + raise RuntimeError("an epoch omitted or duplicated positive support") + if ( + float(cfg.alpha) > 0.0 + and negatives + and any(not row["negative_coverage"]["exact_once"] for row in epoch_rows) + ): + raise RuntimeError("an epoch omitted or duplicated negative support") + summary_keys = ( + "positive_loss", "negative_loss", "positive_norm", "negative_norm", + "rho", "gradient_cosine", + ) + return dict( + alpha=float(cfg.alpha), replay_epochs=int(cfg.replay_epochs), + optimizer_steps=optimizer_steps, + positive_eligible=len(positives), negative_eligible=len(negatives), + positive_total_visits=sum( + row["positive_coverage"]["visited"] for row in epoch_rows + ), + negative_total_visits=sum( + row["negative_coverage"]["visited"] for row in epoch_rows + ), + negative_used_for_training=bool(float(cfg.alpha) > 0.0 and negatives), + exact_complete_replay=True, + fixed_probe=dict(before=fixed_probe_before, after=fixed_probe_after), + module_relative_parameter_drift=_module_relative_drift( + modules_before, modules_after, + ), + visual_encoder_sha_before=encoder_before, + visual_encoder_sha_after=encoder_after, + summaries={key: _numeric_summary(epoch_rows, key) for key in summary_keys}, + epochs=epoch_rows, + ) + + +def run(checkpoint, outdir, cfg, *, device): + cfg.validate() + checkpoint = os.path.abspath(checkpoint) + outdir = os.path.abspath(outdir) + if not os.path.isfile(checkpoint): + raise FileNotFoundError(checkpoint) + checkpoint_sha = BS.sha256_file(checkpoint) + if checkpoint_sha != EXPECTED_CHECKPOINT_SHA256: + raise ValueError( + f"checkpoint SHA mismatch: expected {EXPECTED_CHECKPOINT_SHA256}, got {checkpoint_sha}" + ) + if os.path.exists(outdir): + raise FileExistsError(f"refusing to reuse output directory: {outdir}") + os.makedirs(outdir) + environment = SS.scene_profile(cfg.scene_profile) + policy, _ = GPS.load_sfm_policy(checkpoint, device=device) + frozen_parameters = BS.configure_expansion_trainability(policy) + visual_encoder_sha = BS.module_sha256(policy.enc_grid) + optimizer = torch.optim.Adam( + [parameter for parameter in policy.parameters() if parameter.requires_grad], + lr=cfg.lr, + ) + recent = BS.RecentRounds(os.path.join(outdir, "round_shards"), cfg.W) + proposal_generator = torch.Generator(device=device).manual_seed(cfg.seed) + BX._save_checkpoint(policy, os.path.join(outdir, "round_00.pt"), dict( + round=0, experiment=cfg.arm_name, source_checkpoint=checkpoint, + source_sha256=checkpoint_sha, encoder_sha256=visual_encoder_sha, + recipe=asdict(cfg), + )) + history = [] + with ProcessPoolExecutor(max_workers=cfg.verifier_workers) as executor: + for round_i in range(1, cfg.rounds + 1): + round_start = time.perf_counter() + scenarios = SP.expansion_scenarios(round_i, smoke=False) + replicas = [ + BX.Replica( + scenario_id, gamma, n_ped=environment["n_ped"], + ped_speed_range=tuple(environment["ped_speed_range"]), + ) + for scenario_id in scenarios for gamma in SP.GAMMAS + ] + if len(replicas) != 56: + raise RuntimeError("isolated experiment requires 56 macro-round replicas") + policy.eval() + phi_policy = copy.deepcopy(policy).eval() + for parameter in phi_policy.parameters(): + parameter.requires_grad_(False) + gp, gp_ids = BX.gp_from_recent( + phi_policy, recent, ell=ELL, cap=CAP, lam=cfg.gp_lam, + phi_s=cfg.phi_s, device=device, seed=cfg.seed + round_i * 101, + ) + beta, calibrated_ess = BX._initial_beta( + phi_policy, gp, replicas, cfg, device, cfg.seed + round_i * 1009, + ) + shard = BS.RoundShard(round_i) + gather = BX.gather_macro_round( + policy, phi_policy, gp, beta, replicas, cfg, shard, device, + executor, proposal_generator, + ) + gather.pop("traces", None) + shard_manifest = recent.append_and_save(shard) + replay_start = time.perf_counter() + replay = repeat_complete_replay( + policy, optimizer, recent, cfg, device=device, round_i=round_i, + ) + gather["timers"]["replay"] = time.perf_counter() - replay_start + if BS.module_sha256(policy.enc_grid) != visual_encoder_sha: + raise RuntimeError("visual encoder SHA changed") + checkpoint_path = os.path.join(outdir, f"round_{round_i:02d}.pt") + BX._save_checkpoint(policy, checkpoint_path, dict( + round=round_i, experiment=cfg.arm_name, + source_checkpoint=checkpoint, source_sha256=checkpoint_sha, + encoder_sha256=visual_encoder_sha, recipe=asdict(cfg), + ell=ELL, cap=CAP, beta=float(beta), + )) + record = dict( + round=round_i, experiment=cfg.arm_name, + scenarios=list(scenarios), environment=environment, + beta=float(beta), calibrated_ess_over_K=float(calibrated_ess), + verifier=SM.verifier_manifest(), gp_buffer_ids=gp_ids, + gp=gp.diagnostics(), gather=gather, replay=replay, + shard=shard_manifest, checkpoint=os.path.abspath(checkpoint_path), + checkpoint_sha256=BS.sha256_file(checkpoint_path), + wall_seconds=time.perf_counter() - round_start, + ) + history.append(record) + with open(os.path.join(outdir, "metrics.jsonl"), "a") as stream: + stream.write(json.dumps(record) + "\n") + print(json.dumps({ + "round": round_i, "experiment": cfg.arm_name, + "beta": float(beta), "wall_seconds": record["wall_seconds"], + }), flush=True) + manifest = dict( + status="R2_ALPHA_REPLAY_COMPLETE", experiment=cfg.arm_name, + recipe=asdict(cfg), constants=dict( + ell=ELL, cap=CAP, gp_lambda=GP_LAMBDA, + expected_checkpoint_sha256=EXPECTED_CHECKPOINT_SHA256, + ), + source_checkpoint=checkpoint, source_checkpoint_sha256=checkpoint_sha, + environment=environment, frozen_parameters=frozen_parameters, + visual_encoder_sha=visual_encoder_sha, history=history, + ) + complete_path = os.path.join(outdir, "COMPLETE.json") + temporary_complete = complete_path + ".tmp" + with open(temporary_complete, "w") as stream: + json.dump(manifest, stream, indent=2) + os.replace(temporary_complete, complete_path) + return manifest + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--checkpoint", required=True) + parser.add_argument("--outdir", required=True) + parser.add_argument("--alpha", type=float, choices=ALPHAS, required=True) + parser.add_argument( + "--replay-epochs", type=int, choices=REPLAY_EPOCHS, required=True, + ) + parser.add_argument("--device", default="cuda") + parser.add_argument("--seed", type=int, default=20260723) + parser.add_argument("--verifier-workers", type=int, default=32) + args = parser.parse_args() + cfg = ExperimentConfig( + alpha=args.alpha, replay_epochs=args.replay_epochs, + seed=args.seed, verifier_workers=args.verifier_workers, + ) + run(args.checkpoint, args.outdir, cfg, device=args.device) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/sfm_b1_r2_eval.py b/overnight_run_07_12_sfm/sfm_b1_r2_eval.py new file mode 100644 index 0000000..b95195f --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_r2_eval.py @@ -0,0 +1,833 @@ +"""Canonical raw temperature-one evaluation for SFM rounds 0, 1, and 2. + +This module is intentionally independent of expansion-time acquisition. It +uses one fixed M=50/scenario/gamma bank and one fixed latent-noise bank for +every checkpoint, samples one raw flow window per context at temperature one, +and executes its first action. No RBF tilt, verifier selection, fallback, +guidance, or temperature search is present. + +The archived double-shift M100 baseline is carried only as a labeled reference +in the result and figure footer. It is never inserted as an evaluation point +and is never used to alter a measured M50 value. +""" +from __future__ import annotations + +import argparse +from concurrent.futures import ProcessPoolExecutor +from dataclasses import dataclass, field +import hashlib +import json +import math +import multiprocessing as mp +import os +import re +from typing import Any + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import torch + +import _paths # noqa: F401 +import grid_feats as GF +import grid_policy_sfm as GPS +import sfm_b1_eval as BE +import sfm_hp_history as HH +import sfm_metrics2 as SM +import sfm_protocol as SP +import sfm_scene as SS + + +VERSION = "sfm_b1_r2_raw_m50_v1" +M_PER_GAMMA = 50 +T = int(SP.T) +H = int(SP.H) +NFE = 8 +TEMPERATURE = 1.0 +DEFAULT_EP0 = 260_000 +DEFAULT_NOISE_SEED = 2_026_072_3 + +# Historical reference only. These values were measured with M=100/gamma on +# scenarios 250000:250099 and are not a target for the disjoint M50 run. +ARCHIVED_M100_REFERENCE = { + "role": "separate_historical_reference_not_a_curve_point", + "scene_profile": "double_density_velocity_ood", + "ep0": 250_000, + "M_per_gamma": 100, + "checkpoint_sha256": ( + "1b5179c935d3eeff8824967d707d64cc9bab273949ee1f0e4f190172bab1b215" + ), + "source_commit": "ca7f0d718f8d70cf74833b1c75157caf7f1b13f2", + "SR": 0.7000000000, + "CR": 0.3000000000, + "successful_clearance": 0.1310315136398588, + "successful_time_to_goal": 8.692857142857143, + "note": ( + "Archived raw temp=1/NFE=8 M100 result. It is reported separately and " + "must never be substituted for an independently measured M50 value." + ), +} + + +def _sha256_file(path: str | os.PathLike[str]) -> str: + digest = hashlib.sha256() + with open(path, "rb") as stream: + for chunk in iter(lambda: stream.read(1 << 20), b""): + digest.update(chunk) + return digest.hexdigest() + + +def _sha256_json(payload: Any) -> str: + encoded = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode() + return hashlib.sha256(encoded).hexdigest() + + +def _write_json(path: str | os.PathLike[str], payload: Any) -> None: + path = os.path.abspath(os.fspath(path)) + os.makedirs(os.path.dirname(path), exist_ok=True) + temporary = path + ".tmp" + with open(temporary, "w") as stream: + json.dump(payload, stream, indent=2, allow_nan=False) + os.replace(temporary, path) + + +def _checkpoint_specs(checkpoints: list[str], labels: list[str]) -> list[dict]: + if len(checkpoints) != len(labels) or not checkpoints: + raise ValueError("--checkpoints and --labels must have the same nonzero length") + if len(labels) != len(set(labels)): + raise ValueError("checkpoint labels must be unique") + specs = [] + for checkpoint, label in zip(checkpoints, labels): + match = re.fullmatch(r"r([0-9]+)", str(label)) + if match is None: + raise ValueError(f"checkpoint label {label!r} must have form r0, r1, ...") + path = os.path.abspath(checkpoint) + if not os.path.isfile(path): + raise FileNotFoundError(path) + specs.append(dict(label=str(label), round=int(match.group(1)), checkpoint=path)) + rounds = [spec["round"] for spec in specs] + if rounds != sorted(rounds) or len(rounds) != len(set(rounds)): + raise ValueError("checkpoint labels must be unique and increasing") + return specs + + +def _assert_disjoint_from_archive(ep0: int) -> None: + current = set(range(int(ep0), int(ep0) + M_PER_GAMMA)) + archived = set(range( + int(ARCHIVED_M100_REFERENCE["ep0"]), + int(ARCHIVED_M100_REFERENCE["ep0"]) + + int(ARCHIVED_M100_REFERENCE["M_per_gamma"]), + )) + if current & archived: + raise ValueError( + "the M50 qualification bank must remain disjoint from the archived M100 bank" + ) + + +def _noise_bank(*, ep0: int, d: int, seed: int) -> tuple[np.ndarray, dict]: + contract = { + "version": VERSION, + "ep0": int(ep0), + "M_per_gamma": M_PER_GAMMA, + "gammas": list(map(float, SP.GAMMAS)), + "T": T, + "d": int(d), + "seed": int(seed), + "temperature": TEMPERATURE, + "NFE": NFE, + } + generator = np.random.default_rng(int(seed)) + values = generator.standard_normal( + (len(SP.GAMMAS), M_PER_GAMMA, T, int(d)), dtype=np.float32, + ) + metadata = { + **contract, + "dtype": "float32", + "shape": list(values.shape), + "sha256": hashlib.sha256(values.tobytes(order="C")).hexdigest(), + "CRN": ( + "same (gamma,scenario,step) latent across checkpoints; paired scenario " + "IDs across gamma, independent latent slices across gamma" + ), + } + return values, metadata + + +@dataclass +class _Episode: + gamma_index: int + rollout_index: int + episode: int + gamma: float + humans: list + state: np.ndarray = field(default_factory=lambda: np.zeros(4, np.float32)) + history: HH.HpHistory = field(default_factory=HH.HpHistory) + controls: list[np.ndarray] = field(default_factory=list) + states: list[np.ndarray] = field( + default_factory=lambda: [np.zeros(4, np.float32)] + ) + context_states: list[np.ndarray] = field(default_factory=list) + planned_controls: list[np.ndarray] = field(default_factory=list) + ped_xy: list[np.ndarray] = field(default_factory=list) + ped_vel: list[np.ndarray] = field(default_factory=list) + status: str | None = None + minimum_clearance: float = float("inf") + + +def _clearance(state: np.ndarray, ped_xy: np.ndarray) -> float: + if not len(ped_xy): + return float("inf") + return float( + np.linalg.norm(ped_xy - state[:2][None], axis=1).min() - SS.R_PED + ) + + +def _terminal_check(episode: _Episode, ped_xy: np.ndarray) -> bool: + clearance = _clearance(episode.state, ped_xy) + episode.minimum_clearance = min(episode.minimum_clearance, clearance) + if clearance < 0.0: + episode.status = "collision" + elif float(np.linalg.norm(episode.state[:2] - SS.GOAL)) < 0.5: + episode.status = "success" + return episode.status is not None + + +@torch.no_grad() +def run_batched_raw( + policy, + *, + scene_profile: str, + ep0: int, + noise: np.ndarray, + device: str, +) -> list[dict]: + """Evaluate all 7xM cells with one flow batch per closed-loop tick.""" + environment = SS.scene_profile(scene_profile) + expected = (len(SP.GAMMAS), M_PER_GAMMA, T, int(policy.d)) + if tuple(noise.shape) != expected or noise.dtype != np.float32: + raise ValueError(f"noise bank {noise.shape}/{noise.dtype} != {expected}/float32") + episodes = [ + _Episode( + gamma_index=gamma_index, + rollout_index=rollout_index, + episode=int(ep0) + rollout_index, + gamma=float(gamma), + humans=SS.make_humans( + int(ep0) + rollout_index, + 0, + environment["n_ped"], + tuple(environment["ped_speed_range"]), + ), + ) + for gamma_index, gamma in enumerate(SP.GAMMAS) + for rollout_index in range(M_PER_GAMMA) + ] + + for step in range(T): + active, hp10, lows, histories, latents = [], [], [], [], [] + for episode in episodes: + if episode.status is not None: + continue + ped_xy, ped_vel = SS.collect_humans(episode.humans) + if _terminal_check(episode, ped_xy): + continue + obstacles = np.concatenate([ + ped_xy, + np.full((len(ped_xy), 1), SS.R_PED, np.float32), + ], axis=1) + raw_grid = torch.as_tensor( + GF.axis_grid( + episode.state[:2], + obstacles, + 0.0, + R=SS.R_SENSE, + sensing=SS.R_SENSE, + ) + ) + active.append((episode, ped_xy.copy(), ped_vel.copy())) + hp10.append(episode.history.append(raw_grid)) + lows.append(torch.as_tensor( + GF.low5(episode.state, SS.GOAL, episode.gamma) + )) + histories.append(torch.as_tensor(GF.hist_pad( + np.asarray(episode.controls[-16:]) + if episode.controls else np.zeros((0, 2)), + 16, + ))) + latents.append(noise[ + episode.gamma_index, + episode.rollout_index, + step, + ]) + if not active: + break + + hp10_tensor = torch.stack(hp10).to(device) + low_tensor = torch.stack(lows).to(device) + history_tensor = torch.stack(histories).to(device) + context = policy.ctx_from(hp10_tensor, low_tensor, history_tensor) + windows = BE.integrate_latents( + policy, + torch.as_tensor(np.asarray(latents), device=device), + context, + nfe=NFE, + ).reshape(len(active), H, 2) + windows = windows.detach().cpu().numpy().astype(np.float32) + + for (episode, ped_xy, ped_vel), window in zip(active, windows): + if tuple(window.shape) != (H, 2): + raise RuntimeError(f"generated plan {window.shape} != {(H, 2)}") + episode.context_states.append(episode.state.copy()) + episode.planned_controls.append(window.copy()) + episode.ped_xy.append(ped_xy) + episode.ped_vel.append(ped_vel) + action = window[0].copy() + episode.controls.append(action) + episode.state = BE._step(episode.state, action) + episode.states.append(episode.state.copy()) + SS.advance_humans(episode.humans, episode.state) + + rows = [] + for episode in episodes: + if episode.status is None: + ped_xy, _ = SS.collect_humans(episode.humans) + if not _terminal_check(episode, ped_xy): + episode.status = "timeout" + success = episode.status == "success" + rows.append({ + "episode": episode.episode, + "gamma": episode.gamma, + "status": episode.status, + "success": success, + "collision": episode.status == "collision", + "timeout": episode.status == "timeout", + "steps": len(episode.controls), + "time_to_goal": ( + len(episode.controls) * SS.DT if success else None + ), + "min_clearance": float(episode.minimum_clearance), + "successful_clearance": ( + float(episode.minimum_clearance) if success else None + ), + "states": np.asarray(episode.states, np.float32), + "context_states": np.asarray(episode.context_states, np.float32), + "planned_controls": np.asarray(episode.planned_controls, np.float32), + "ped_xy": np.asarray(episode.ped_xy, np.float32), + "ped_vel": np.asarray(episode.ped_vel, np.float32), + }) + return rows + + +def _verify_episode(row: dict) -> dict: + n_steps = int(row["steps"]) + states = np.asarray(row["states"], np.float32) + context_states = np.asarray(row["context_states"], np.float32) + planned_controls = np.asarray(row["planned_controls"], np.float32) + ped_xy = np.asarray(row["ped_xy"], np.float32) + ped_vel = np.asarray(row["ped_vel"], np.float32) + expected_lengths = ( + len(states) == n_steps + 1 + and len(context_states) == n_steps + and len(planned_controls) == n_steps + and len(ped_xy) == n_steps + and len(ped_vel) == n_steps + ) + if not expected_lengths or ( + n_steps and tuple(planned_controls.shape[1:]) != (H, 2) + ): + return {"v_safe": False, "verifier_errors": 1, "certified_windows": 0} + + physical_safe = ( + not bool(row["collision"]) + and SM.taskspace_ok(states[:, :2]) + and n_steps > 0 + ) + if not physical_safe: + return {"v_safe": False, "verifier_errors": 0, "certified_windows": 0} + + certified_windows = 0 + for state, controls, current_xy, current_vel in zip( + context_states, planned_controls, ped_xy, ped_vel + ): + result = SM.verify_query( + state, controls, current_xy, current_vel, float(row["gamma"]) + ) + if not result.get("resolved", False): + return { + "v_safe": False, + "verifier_errors": 1, + "certified_windows": certified_windows, + } + if not result.get("full_h", False) or int(result.get("terminal_step", -1)) != H: + return { + "v_safe": False, + "verifier_errors": 1, + "certified_windows": certified_windows, + } + certified_windows += 1 + if not bool(result["y"]): + return { + "v_safe": False, + "verifier_errors": 0, + "certified_windows": certified_windows, + } + return { + "v_safe": True, + "verifier_errors": 0, + "certified_windows": certified_windows, + } + + +def _attach_validity(rows: list[dict], executor) -> list[dict]: + futures = [executor.submit(_verify_episode, row) for row in rows] + compact = [] + omitted = { + "states", "context_states", "planned_controls", "ped_xy", "ped_vel", + } + for row, future in zip(rows, futures): + value = {key: item for key, item in row.items() if key not in omitted} + value.update(future.result()) + compact.append(value) + return compact + + +def _cluster_bootstrap_interval( + rows: list[dict], + key: str, + *, + seed: int, + draws: int = 2_000, +) -> list[float | None]: + episode_ids = sorted({int(row["episode"]) for row in rows}) + sums, counts = [], [] + for episode in episode_ids: + values = [row.get(key) for row in rows if int(row["episode"]) == episode] + finite = [ + float(value) for value in values + if value is not None and math.isfinite(float(value)) + ] + sums.append(sum(finite)) + counts.append(len(finite)) + if not episode_ids or not sum(counts): + return [None, None] + generator = np.random.default_rng(int(seed)) + indices = generator.integers( + 0, len(episode_ids), size=(int(draws), len(episode_ids)) + ) + numerator = np.asarray(sums, float)[indices].sum(axis=1) + denominator = np.asarray(counts, float)[indices].sum(axis=1) + samples = numerator[denominator > 0] / denominator[denominator > 0] + if not len(samples): + return [None, None] + return list(map(float, np.quantile(samples, [.025, .975]))) + + +def _summarize_one(rows: list[dict], seed: int) -> dict: + n = len(rows) + if n < 1: + raise ValueError("cannot summarize an empty cell") + successes = sum(bool(row["success"]) for row in rows) + collisions = sum(bool(row["collision"]) for row in rows) + timeouts = sum(bool(row["timeout"]) for row in rows) + valid = sum(bool(row["v_safe"]) for row in rows) + if successes + collisions + timeouts != n: + raise RuntimeError("success, collision, and timeout must partition a cell") + return { + "n": n, + "SR": successes / n, + "SR_wilson95": BE.wilson(successes, n), + "CR": collisions / n, + "CR_wilson95": BE.wilson(collisions, n), + "timeout": timeouts / n, + "timeout_wilson95": BE.wilson(timeouts, n), + "V_safe": valid / n, + "V_safe_wilson95": BE.wilson(valid, n), + "successful_clearance": BE.bootstrap_mean( + [row["successful_clearance"] for row in rows], seed=seed + ), + "successful_time_to_goal": BE.bootstrap_mean( + [row["time_to_goal"] for row in rows], seed=seed + 1 + ), + "verifier_errors": sum(int(row["verifier_errors"]) for row in rows), + "certified_windows": sum(int(row["certified_windows"]) for row in rows), + } + + +def summarize(rows: list[dict], *, seed: int) -> dict: + per_gamma = { + str(gamma): _summarize_one( + [row for row in rows if float(row["gamma"]) == float(gamma)], + seed + index * 10, + ) + for index, gamma in enumerate(SP.GAMMAS) + } + pooled = _summarize_one(rows, seed + 100) + for metric, key in ( + ("SR", "success"), + ("CR", "collision"), + ("timeout", "timeout"), + ("V_safe", "v_safe"), + ): + pooled[f"{metric}_cluster_bootstrap95"] = _cluster_bootstrap_interval( + rows, key, seed=seed + 200 + len(metric) + ) + pooled["successful_clearance"]["cluster_bootstrap95"] = ( + _cluster_bootstrap_interval( + rows, "successful_clearance", seed=seed + 300 + ) + ) + pooled["successful_time_to_goal"]["cluster_bootstrap95"] = ( + _cluster_bootstrap_interval( + rows, "time_to_goal", seed=seed + 301 + ) + ) + pooled["ci_method"] = ( + "scenario-cluster bootstrap across seven paired gamma rows" + ) + return {"pooled": pooled, "per_gamma": per_gamma} + + +def _assert_zero_verifier_errors(summary: dict) -> None: + cells = [summary["pooled"], *summary["per_gamma"].values()] + if any(int(cell["verifier_errors"]) != 0 for cell in cells): + raise RuntimeError("evaluation contains verifier errors") + + +def _cell_key( + *, + checkpoint_sha256: str, + scene_profile: str, + ep0: int, + noise_meta: dict, +) -> str: + return _sha256_json({ + "version": VERSION, + "evaluator_sha256": _sha256_file(__file__), + "checkpoint_sha256": checkpoint_sha256, + "scene_profile": scene_profile, + "ep0": int(ep0), + "M_per_gamma": M_PER_GAMMA, + "noise_bank": noise_meta, + "temperature": TEMPERATURE, + "NFE": NFE, + "T": T, + "H": H, + "verifier": SM.verifier_manifest(), + }) + + +def _evaluate_checkpoint( + checkpoint: str, + *, + scene_profile: str, + ep0: int, + noise: np.ndarray, + noise_meta: dict, + device: str, + cache_dir: str, + executor, +) -> dict: + checkpoint_sha = _sha256_file(checkpoint) + key = _cell_key( + checkpoint_sha256=checkpoint_sha, + scene_profile=scene_profile, + ep0=ep0, + noise_meta=noise_meta, + ) + cache_path = os.path.join( + cache_dir, f"cell_{checkpoint_sha[:12]}_{key[:12]}.json" + ) + if os.path.isfile(cache_path): + with open(cache_path) as stream: + payload = json.load(stream) + if ( + payload.get("status") != "SFM_B1_R2_RAW_CELL_COMPLETE" + or payload.get("cell_key") != key + ): + raise RuntimeError(f"stale evaluation cache: {cache_path}") + _assert_zero_verifier_errors(payload["summary"]) + return payload + + policy, _ = GPS.load_sfm_policy(checkpoint, device=device) + policy.eval() + if int(policy.d) != int(noise.shape[-1]): + raise ValueError("checkpoint latent dimension does not match the fixed noise bank") + rows = run_batched_raw( + policy, + scene_profile=scene_profile, + ep0=ep0, + noise=noise, + device=device, + ) + del policy + if str(device).startswith("cuda"): + torch.cuda.empty_cache() + compact = _attach_validity(rows, executor) + summary = summarize( + compact, + seed=int(ep0) + int(checkpoint_sha[:8], 16) % 100_000, + ) + _assert_zero_verifier_errors(summary) + payload = { + "status": "SFM_B1_R2_RAW_CELL_COMPLETE", + "cell_key": key, + "checkpoint": os.path.abspath(checkpoint), + "checkpoint_sha256": checkpoint_sha, + "scene_profile": scene_profile, + "ep0": int(ep0), + "M_per_gamma": M_PER_GAMMA, + "summary": summary, + "rows": compact, + "metric_semantics": { + "policy": ( + "canonical unguided raw flow, temperature=1, NFE=8, one " + "generated H=10 window per context, execute first action" + ), + "V_safe": ( + "episode is physically collision/task-space safe and every " + "generated plan at every executed context passes the exact " + "full-H=10 moving-pedestrian verifier" + ), + "clearance": ( + "mean of each successful trajectory's minimum pedestrian " + "clearance; failures are excluded" + ), + "time": "successful trajectories only", + "outcome_partition": "SR + CR + timeout = 1", + }, + } + _write_json(cache_path, payload) + return payload + + +def _metric_value(cell: dict, metric: str) -> float: + if metric in ("CR", "V_safe"): + return float(cell[metric]) + key = ( + "successful_clearance" + if metric == "clearance" + else "successful_time_to_goal" + ) + value = cell[key]["mean"] + return float("nan") if value is None else float(value) + + +def _pooled_interval(cell: dict, metric: str) -> list[float]: + if metric in ("CR", "V_safe"): + value = cell[f"{metric}_cluster_bootstrap95"] + else: + key = ( + "successful_clearance" + if metric == "clearance" + else "successful_time_to_goal" + ) + value = cell[key]["cluster_bootstrap95"] + return [ + float("nan") if item is None else float(item) + for item in value + ] + + +def render(records: list[dict], output_dir: str) -> list[str]: + """Render the four requested metrics in the B1 paper-curve style.""" + colors = plt.get_cmap("plasma")( + np.linspace(0.08, 0.92, len(SP.GAMMAS)) + ) + specs = ( + ("CR", "Collision rate"), + ("V_safe", r"$V_{\mathrm{safe}}$"), + ("clearance", "Successful min. clearance [m]"), + ("time", "Successful time-to-goal [s]"), + ) + rounds = [int(record["round"]) for record in records] + plt.rcParams.update({ + "font.family": "serif", + "mathtext.fontset": "cm", + "font.serif": ["cmr10", "Computer Modern Roman", "DejaVu Serif"], + "axes.unicode_minus": False, + "axes.formatter.use_mathtext": True, + "axes.titlesize": 18, + "axes.labelsize": 16, + "xtick.labelsize": 13, + "ytick.labelsize": 13, + }) + figure, axes = plt.subplots(2, 2, figsize=(14.5, 9)) + for axis, (metric, title) in zip(axes.flat, specs): + for gamma, color in zip(SP.GAMMAS, colors): + cells = [ + record["cell"]["summary"]["per_gamma"][str(gamma)] + for record in records + ] + axis.plot( + rounds, + [_metric_value(cell, metric) for cell in cells], + color=color, + lw=1.5, + marker="o", + ms=5, + alpha=0.72, + label=rf"$\gamma={gamma:g}$", + ) + pooled = [ + record["cell"]["summary"]["pooled"] for record in records + ] + values = [_metric_value(cell, metric) for cell in pooled] + intervals = [_pooled_interval(cell, metric) for cell in pooled] + axis.plot( + rounds, + values, + color="black", + lw=3.0, + marker="o", + ms=6, + label=r"pooled ($7\gamma$)", + zorder=4, + ) + axis.fill_between( + rounds, + [value[0] for value in intervals], + [value[1] for value in intervals], + color="black", + alpha=0.11, + lw=0, + zorder=1, + ) + axis.set_title(title, pad=8) + axis.set_xlabel("expansion round") + axis.set_xticks(rounds) + axis.grid(alpha=0.25) + axis.set_xlim(min(rounds) - 0.15, max(rounds) + 0.15) + if metric in ("CR", "V_safe"): + axis.set_ylim(-0.03, 1.03) + + handles, labels = axes[0, 0].get_legend_handles_labels() + figure.legend( + handles, + labels, + loc="upper center", + ncol=4, + frameon=False, + bbox_to_anchor=(0.5, 0.995), + ) + reference = ARCHIVED_M100_REFERENCE + figure.text( + 0.5, + 0.012, + ( + "Separate archived M100 reference (not plotted): " + f"SR {reference['SR']:.3f}, CR {reference['CR']:.3f}, " + f"successful clearance {reference['successful_clearance']:.3f} m, " + f"successful time {reference['successful_time_to_goal']:.3f} s." + ), + ha="center", + va="bottom", + fontsize=11, + color="0.35", + ) + figure.tight_layout(rect=(0.02, 0.055, 0.98, 0.91)) + os.makedirs(output_dir, exist_ok=True) + outputs = [] + for suffix in ("png", "pdf"): + path = os.path.join(output_dir, f"raw_m50_r0_r2_curves.{suffix}") + figure.savefig(path, dpi=300, bbox_inches="tight") + outputs.append(path) + plt.close(figure) + return outputs + + +def run(args) -> dict: + specs = _checkpoint_specs(args.checkpoints, args.labels) + _assert_disjoint_from_archive(args.ep0) + output_dir = os.path.abspath(args.output_dir) + cache_dir = os.path.abspath(args.cache_dir or os.path.join(output_dir, "cache")) + os.makedirs(output_dir, exist_ok=True) + os.makedirs(cache_dir, exist_ok=True) + + probe, _ = GPS.load_sfm_policy(specs[0]["checkpoint"], device="cpu") + noise, noise_meta = _noise_bank( + ep0=args.ep0, d=int(probe.d), seed=args.noise_seed + ) + del probe + records = [] + context = mp.get_context("spawn") + with ProcessPoolExecutor( + max_workers=int(args.workers), mp_context=context + ) as executor: + for spec in specs: + cell = _evaluate_checkpoint( + spec["checkpoint"], + scene_profile=args.scene_profile, + ep0=args.ep0, + noise=noise, + noise_meta=noise_meta, + device=args.device, + cache_dir=cache_dir, + executor=executor, + ) + records.append({ + "label": spec["label"], + "round": spec["round"], + "cell": cell, + }) + + outputs = render(records, output_dir) + result = { + "status": "SFM_B1_R2_RAW_M50_COMPLETE", + "version": VERSION, + "scene_profile": args.scene_profile, + "environment": SS.scene_profile(args.scene_profile), + "bank": { + "ep0": int(args.ep0), + "M_per_gamma": M_PER_GAMMA, + "scenario_ids": list(range( + int(args.ep0), int(args.ep0) + M_PER_GAMMA + )), + "same_scenario_ids_for_every_gamma": True, + "disjoint_from_archived_M100": True, + }, + "noise_bank": noise_meta, + "records": records, + "archived_M100_reference": ARCHIVED_M100_REFERENCE, + "reference_policy": ( + "The archived M100 result is provenance only. No measured M50 " + "value is replaced, shifted, selected, or calibrated against it." + ), + "outputs": outputs, + } + result_path = os.path.join(output_dir, "raw_m50_r0_r2_metrics.json") + _write_json(result_path, result) + result["metrics_json"] = result_path + return result + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--checkpoints", nargs="+", required=True) + parser.add_argument("--labels", nargs="+", required=True) + parser.add_argument( + "--scene-profile", + default="double_density_velocity_ood", + choices=SS.SCIENTIFIC_EVAL_PROFILES, + ) + parser.add_argument("--ep0", type=int, default=DEFAULT_EP0) + parser.add_argument("--noise-seed", type=int, default=DEFAULT_NOISE_SEED) + parser.add_argument("--device", default="cuda") + parser.add_argument("--workers", type=int, default=32) + parser.add_argument("--cache-dir") + parser.add_argument("--output-dir", required=True) + return parser + + +def main(argv=None) -> int: + args = build_parser().parse_args(argv) + result = run(args) + print(result["metrics_json"]) + for path in result["outputs"]: + print(path) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/overnight_run_07_12_sfm/sfm_b1_selector_compare_viz.py b/overnight_run_07_12_sfm/sfm_b1_selector_compare_viz.py new file mode 100644 index 0000000..6196200 --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_selector_compare_viz.py @@ -0,0 +1,221 @@ +"""Compare pretrained max-margin and native-SafeMPPI-cost data acquisition.""" +from __future__ import annotations + +import argparse +import json +import os + +import matplotlib +matplotlib.use("Agg") +import matplotlib.animation as animation +import matplotlib.pyplot as plt +import torch + +import sfm_b1_branch_compare_viz as BC +import sfm_b1_d_branch_viz as DB +import sfm_b1_full_episode_viz as FV + + +STATUS = "SFM_B1_SELECTOR_COMPARISON_COMPLETE" +SELECTORS = ("margin", "safemppi_cost", "balanced_rank") + + +def _load(path, selector): + bundle = torch.load(path, map_location="cpu", weights_only=False) + if bundle.get("status") != "SFM_B1_FULL_EPISODE_LABEL_AUDIT_COMPLETE": + raise ValueError(f"not a completed branch audit: {path}") + observed = bundle.get("protocol", {}).get("selector") + if observed != selector: + raise ValueError( + f"expected selector={selector}, observed {observed} in {path}" + ) + return bundle + + +def _validate_bundles(bundles): + reference = bundles[0][1] + for _, bundle in bundles[1:]: + for key in ( + "scenarios", "gammas", "environment", "sample_seed", "audit_seed", + ): + if reference[key] != bundle[key]: + raise ValueError( + f"selector comparison contract differs at {key}" + ) + if ( + reference.get("checkpoint_sha256") + != bundle.get("checkpoint_sha256") + ): + raise ValueError( + "selector comparison requires one pretrained checkpoint" + ) + if len(reference["scenarios"]) != 3 or len(reference["gammas"]) != 7: + raise ValueError("selector comparison requires 3 episodes x 7 gammas") + + +def _layout(bundles): + scenarios = tuple(map(int, bundles[0][1]["scenarios"])) + gammas = tuple(map(float, bundles[0][1]["gammas"])) + rows = len(bundles) * len(scenarios) + figure, axes = plt.subplots(rows, 7, figsize=(23.5, 3.0 * rows)) + figure.subplots_adjust( + left=.055, right=.82, bottom=.025, top=.96, wspace=.025, hspace=.04, + ) + for column, gamma in enumerate(gammas): + figure.text( + .055 + (.765 / 7) * (column + .5), .975, + f"$\\gamma={gamma:g}$", ha="center", va="center", fontsize=10, + ) + indices = {} + for selector_index, (label, bundle) in enumerate(bundles): + index = FV._index(bundle["traces"]) + indices[label] = index + for scenario_index, scenario in enumerate(scenarios): + row = selector_index * len(scenarios) + scenario_index + figure.text( + .018, .96 - (.935 / rows) * (row + .5), + f"{label}\nepisode {scenario}", + ha="center", va="center", rotation=90, fontsize=8, + ) + return figure, axes, scenarios, gammas, bundles, indices + + +def _draw(axes, scenarios, gammas, bundles, indices, step): + for selector_index, (label, _) in enumerate(bundles): + index = indices[label] + for scenario_index, scenario in enumerate(scenarios): + row = selector_index * len(scenarios) + scenario_index + for column, gamma in enumerate(gammas): + axis = axes[row, column] + axis.clear() + DB.draw_cell( + axis, index[(scenario, round(gamma, 8))], int(step), + branch_line_scale=2.7, + trajectory_linewidth=1.05, + trajectory_marker_size=1.25, + candidate_inset=True, + ) + + +def render( + margin_trace, cost_trace, output_png, output_mp4, output_json, + *, balanced_trace=None, fps=5, frame_stride=2, +): + margin = _load(margin_trace, "margin") + cost = _load(cost_trace, "safemppi_cost") + bundles = [ + ("max one-step margin", margin), + ("SafeMPPI cost", cost), + ] + if balanced_trace is not None: + bundles.append(( + "balanced safety + performance rank", + _load(balanced_trace, "balanced_rank"), + )) + _validate_bundles(bundles) + ( + figure, axes, scenarios, gammas, bundles, indices, + ) = _layout(tuple(bundles)) + + reports = {label: BC.summarize(bundle) for label, bundle in bundles} + maximum = max( + max(max(rows) for rows in index.values()) + for index in indices.values() + ) + frames = list(range(0, maximum + 1, int(frame_stride))) + if frames[-1] != maximum: + frames.append(maximum) + + figure.legend( + handles=DB._legend(), loc="center left", bbox_to_anchor=(.835, .72), + frameon=False, fontsize=8, + ) + summary = [] + for label, report in reports.items(): + summary.extend([ + label, + f"positive D: {report['executed_positive']}/{report['contexts']}", + f"S/C/T: {report['success']}/" + f"{report['collision']}/{report['timeout']}", + "", + ]) + summary.extend([ + "Same pretrained checkpoint, episodes,", + "gammas, proposal-noise contract.", + "Only the admissible-B execution ranking differs.", + "Branches: planned H10 D samples.", + "Thin black: executed first-action path.", + ]) + figure.text( + .835, .47, "\n".join(summary), ha="left", va="top", fontsize=8, + ) + + for path in (output_png, output_mp4, output_json): + os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True) + _draw(axes, scenarios, gammas, bundles, indices, maximum) + figure.savefig(output_png, dpi=165, bbox_inches="tight") + + def update(step): + _draw(axes, scenarios, gammas, bundles, indices, int(step)) + return [] + + movie = animation.FuncAnimation( + figure, update, frames=frames, interval=1000 / int(fps), blit=False, + ) + movie.save( + output_mp4, writer=animation.FFMpegWriter( + fps=int(fps), bitrate=5200, + ), dpi=105, + ) + plt.close(figure) + + report = { + "status": STATUS, + "margin_trace": os.path.abspath(margin_trace), + "safemppi_cost_trace": os.path.abspath(cost_trace), + "balanced_rank_trace": ( + None if balanced_trace is None + else os.path.abspath(balanced_trace) + ), + "checkpoint_sha256": margin.get("checkpoint_sha256"), + "scenarios": list(scenarios), + "gammas": list(gammas), + "comparison": reports, + "controlled_difference": ( + "rank the same SOCP-positive and nominal-Hp-admissible B queries " + "by max one-step Hp margin, minimum frozen native SafeMPPI cost, " + "or the sum of ordinal safety/performance ranks with safety-first " + "tie-breaking" + ), + "frames": frames, + "png": os.path.abspath(output_png), + "mp4": os.path.abspath(output_mp4), + } + temporary = os.path.abspath(output_json) + ".tmp" + with open(temporary, "w") as stream: + json.dump(report, stream, indent=2, allow_nan=False) + os.replace(temporary, os.path.abspath(output_json)) + return report + + +def main(argv=None): + parser = argparse.ArgumentParser() + parser.add_argument("--margin-trace", required=True) + parser.add_argument("--safemppi-cost-trace", required=True) + parser.add_argument("--balanced-rank-trace") + parser.add_argument("--output-png", required=True) + parser.add_argument("--output-mp4", required=True) + parser.add_argument("--output-json", required=True) + parser.add_argument("--fps", type=int, default=5) + parser.add_argument("--frame-stride", type=int, default=2) + args = parser.parse_args(argv) + render( + args.margin_trace, args.safemppi_cost_trace, + args.output_png, args.output_mp4, args.output_json, + balanced_trace=args.balanced_rank_trace, + fps=args.fps, frame_stride=args.frame_stride, + ) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/sfm_b1_teacher_branch_viz.py b/overnight_run_07_12_sfm/sfm_b1_teacher_branch_viz.py new file mode 100644 index 0000000..b2ab11a --- /dev/null +++ b/overnight_run_07_12_sfm/sfm_b1_teacher_branch_viz.py @@ -0,0 +1,435 @@ +"""Render ordinary executed D branches beside a separate MPC-teacher buffer. + +The two inputs must be the exact matching macro-round shard and teacher +buffer. Blue/red are the ordinary exact-verifier labels in D; purple is a +privileged Codex MPC teacher target and is deliberately never presented as a +certificate or as D+. +""" +from __future__ import annotations + +import argparse +from collections import defaultdict +import hashlib +import json +import os + +import matplotlib +matplotlib.use("Agg") +import matplotlib.animation as animation +import matplotlib.pyplot as plt +from matplotlib.lines import Line2D +import numpy as np +import torch + +import _paths # noqa: F401 +import sfm_b1_offline_store as OS +import sfm_b1_viz as BV +import sfm_metrics2 as SM +import sfm_scene as SS + + +TEACHER_STATUS = "SFM_UNVERIFIED_MPC_TEACHER_BUFFER_COMPLETE" +TEACHER_SOURCE = "codex_privileged_sfm_mpc" +D_POS_COLOR = "#0067C5" +D_NEG_COLOR = "#C62828" +TEACHER_COLOR = "#7B2CBF" +CONTEXT_FIELDS = ("scenario_id", "gamma", "step", "state", "ped_xy", "ped_vel") + + +def _sha256(path): + digest = hashlib.sha256() + with open(path, "rb") as stream: + for chunk in iter(lambda: stream.read(1 << 20), b""): + digest.update(chunk) + return digest.hexdigest() + + +def _write_json(path, payload): + os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True) + temporary = os.fspath(path) + ".tmp" + with open(temporary, "w") as stream: + json.dump(payload, stream, indent=2, sort_keys=True) + os.replace(temporary, path) + + +def _context_snapshot(bundle, record): + for key in ("context_snapshot", "context"): + value = record.get(key) + if isinstance(value, dict): + return value + snapshots = bundle.get("context_snapshots") + context_id = int(record["context_id"]) + if isinstance(snapshots, list) and 0 <= context_id < len(snapshots): + return snapshots[context_id] + if isinstance(snapshots, dict): + value = snapshots.get(context_id, snapshots.get(str(context_id))) + if isinstance(value, dict): + return value + if all(field in record for field in CONTEXT_FIELDS): + return record + raise ValueError( + f"teacher {record.get('teacher_id')} has no auditable context snapshot" + ) + + +def _same_array(left, right): + left = np.asarray(left) + right = np.asarray(right) + return left.shape == right.shape and np.array_equal(left, right) + + +def _validate_snapshot(snapshot, context, teacher_id): + scalar_match = ( + int(snapshot["scenario_id"]) == int(context["scenario_id"]) + and round(float(snapshot["gamma"]), 8) + == round(float(context["gamma"]), 8) + and int(snapshot["step"]) == int(context["step"]) + ) + array_match = all( + _same_array(snapshot[field], context[field]) + for field in ("state", "ped_xy", "ped_vel") + ) + if not scalar_match or not array_match: + raise ValueError( + f"teacher {teacher_id} context snapshot does not match round shard" + ) + + +def load_inputs(round_shard_path, teacher_buffer_path): + """Load and fail closed unless every teacher row matches the exact shard.""" + round_shard_path = os.path.abspath(round_shard_path) + teacher_buffer_path = os.path.abspath(teacher_buffer_path) + shard_sha256 = _sha256(round_shard_path) + shard = OS.ExecutedRoundShard.load(round_shard_path) + bundle = torch.load( + teacher_buffer_path, map_location="cpu", weights_only=False, + ) + if bundle.get("status") != TEACHER_STATUS: + raise ValueError("input is not a completed unverified MPC teacher buffer") + if int(bundle.get("version", -1)) != 1: + raise ValueError("unsupported unverified MPC teacher buffer version") + if int(bundle.get("round", -1)) != int(shard.round_i): + raise ValueError("teacher buffer round does not match D round shard") + if bundle.get("round_shard_sha256") != shard_sha256: + raise ValueError("teacher buffer does not authenticate the exact D round shard") + + records = list(bundle.get("records", ())) + seen = set() + by_context = defaultdict(list) + family_counts = defaultdict(int) + for index, record in enumerate(records): + teacher_id = int(record.get("teacher_id", -1)) + if teacher_id < 0 or teacher_id in seen: + raise ValueError("teacher IDs must be unique non-negative integers") + seen.add(teacher_id) + if record.get("source") != TEACHER_SOURCE: + raise ValueError(f"teacher {teacher_id} has an undeclared source") + forbidden = {"y", "query_id", "train_eligible", "x0"} & set(record) + if forbidden: + raise ValueError( + f"teacher {teacher_id} illegally carries safety-label fields " + f"{sorted(forbidden)}" + ) + context_id = int(record["context_id"]) + if not 0 <= context_id < len(shard.contexts): + raise ValueError(f"teacher {teacher_id} references missing context") + context = shard.contexts[context_id] + _validate_snapshot( + _context_snapshot(bundle, record), context, teacher_id, + ) + controls = np.asarray(record["controls"], np.float32) + if ( + controls.shape != (10, 2) + or not np.isfinite(controls).all() + or float(np.max(np.abs(controls))) > float(SS.U_MAX) + 1.0e-6 + ): + raise ValueError( + f"teacher {teacher_id} controls violate the H10/input-limit contract" + ) + normalized = dict(record) + normalized["teacher_id"] = teacher_id + normalized["context_id"] = context_id + normalized["controls"] = controls + by_context[context_id].append(normalized) + candidate_source = record.get("candidate_source") + family = ( + candidate_source.get("family", "unknown") + if isinstance(candidate_source, dict) + else record.get( + "candidate_family", record.get("family", "unknown"), + ) + ) + family_counts[str(family)] += 1 + return shard, bundle, dict(by_context), dict( + round_shard_path=round_shard_path, + round_shard_sha256=shard_sha256, + teacher_buffer_path=teacher_buffer_path, + teacher_buffer_sha256=_sha256(teacher_buffer_path), + teacher_records=len(records), + teacher_contexts=len(by_context), + teacher_families=dict(sorted(family_counts.items())), + ) + + +def _lineages(shard): + output = defaultdict(list) + for window in shard.windows: + context = shard.contexts[int(window["context_id"])] + key = ( + int(context["scenario_id"]), + round(float(context["gamma"]), 8), + ) + output[key].append((int(context["step"]), context, window)) + for key, rows in output.items(): + rows.sort(key=lambda item: item[0]) + if len({step for step, _, _ in rows}) != len(rows): + raise ValueError(f"duplicate D step in lineage {key}") + return dict(output) + + +def _complete_scenarios(lineages): + expected = {round(float(gamma), 8) for gamma in SS.GAMMAS} + found = defaultdict(set) + for scenario, gamma in lineages: + found[int(scenario)].add(gamma) + return tuple( + scenario for scenario in sorted(found) + if found[scenario] == expected + ) + + +def _next_state(state, controls): + state = np.asarray(state, np.float32) + action = np.asarray(controls, np.float32)[0] + next_state = state.copy() + next_state[:2] = ( + state[:2] + SS.DT * state[2:4] + 0.5 * SS.DT ** 2 * action + ) + next_state[2:4] = state[2:4] + SS.DT * action + return next_state + + +def draw_cell(axis, rows, teachers_by_context, through_step): + available = [row for row in rows if row[0] <= int(through_step)] + if not available: + available = [rows[0]] + current_context = available[-1][1] + BV._draw_common(axis, current_context, nominal_levels=False) + axis.plot(SS.GOAL[0], SS.GOAL[1], "*", color="#009E73", ms=8, zorder=12) + + trajectory = [np.asarray(context["state"], float)[:2] + for _, context, _ in available] + last = available[-1] + trajectory.append( + _next_state(last[1]["state"], last[2]["controls"])[:2], + ) + trajectory = np.asarray(trajectory) + + for _, context, window in available: + path = SM.rollout_positions(context["state"], window["controls"]) + color = D_POS_COLOR if int(window["y"]) == 1 else D_NEG_COLOR + axis.plot( + path[:, 0], path[:, 1], color=color, lw=.62, + marker=".", ms=.85, alpha=.36, zorder=3, + ) + for teacher in teachers_by_context.get(int(context["context_id"]), ()): + teacher_path = SM.rollout_positions( + context["state"], teacher["controls"], + ) + axis.plot( + teacher_path[:, 0], teacher_path[:, 1], + color=TEACHER_COLOR, ls="--", lw=1.35, + marker=".", ms=1.25, alpha=.82, zorder=6, + ) + + axis.plot( + trajectory[:, 0], trajectory[:, 1], color="#111111", lw=2.2, + marker=".", ms=1.7, alpha=.97, zorder=10, + ) + axis.set_xticks([]) + axis.set_yticks([]) + return available[-1][0] + + +def _legend(): + return [ + Line2D([], [], color=D_POS_COLOR, lw=1.2, + label=r"ordinary $D^+$ · exact full-H positive"), + Line2D([], [], color=D_NEG_COLOR, lw=1.2, + label=r"ordinary $D^-$ · exact full-H negative"), + Line2D([], [], color="#111111", lw=2.2, + label="ordinary executed first-action trajectory"), + Line2D([], [], color=TEACHER_COLOR, ls="--", lw=1.8, + label=r"$D_{\rm MPC}$ privileged dodge teacher · unverified"), + ] + + +def render( + round_shard_path, teacher_buffer_path, output_png, output_json, *, + scenarios=None, output_mp4=None, fps=5, frame_stride=2, +): + if int(fps) <= 0 or int(frame_stride) <= 0: + raise ValueError("fps and frame_stride must be positive") + shard, teacher_bundle, teachers, provenance = load_inputs( + round_shard_path, teacher_buffer_path, + ) + lineages = _lineages(shard) + scenario_teacher_counts = defaultdict(int) + for context_id, rows in teachers.items(): + scenario = int(shard.contexts[int(context_id)]["scenario_id"]) + scenario_teacher_counts[scenario] += len(rows) + scenarios_were_explicit = scenarios is not None + if scenarios is None: + candidates = _complete_scenarios(lineages) + if len(candidates) < 3: + raise ValueError("round shard has fewer than three complete 7-gamma episodes") + scenarios = tuple(sorted( + candidates, + key=lambda scenario: ( + -scenario_teacher_counts[int(scenario)], int(scenario), + ), + )[:3]) + scenarios = tuple(map(int, scenarios)) + if len(scenarios) != 3 or len(set(scenarios)) != 3: + raise ValueError("renderer requires exactly three distinct scenarios") + gammas = tuple(map(float, SS.GAMMAS)) + missing = [ + (scenario, gamma) + for scenario in scenarios for gamma in gammas + if (scenario, round(gamma, 8)) not in lineages + ] + if missing: + raise ValueError(f"selected 3x7 grid is incomplete: {missing}") + + maximum = max( + rows[-1][0] for key, rows in lineages.items() + if key[0] in scenarios + ) + frames = list(range(0, maximum + 1, int(frame_stride))) + if frames[-1] != maximum: + frames.append(maximum) + figure, axes = plt.subplots(3, 7, figsize=(23.2, 10.1)) + figure.subplots_adjust( + left=.035, right=.805, bottom=.025, top=.94, + wspace=.025, hspace=.04, + ) + for column, gamma in enumerate(gammas): + figure.text( + .035 + (.77 / 7) * (column + .5), .965, + f"$\\gamma={gamma:g}$", ha="center", va="center", fontsize=10, + ) + for row, scenario in enumerate(scenarios): + figure.text( + .012, .94 - (.915 / 3) * (row + .5), + f"episode\n{scenario}", ha="center", va="center", + rotation=90, fontsize=9, + ) + figure.legend( + handles=_legend(), loc="center left", bbox_to_anchor=(.815, .63), + frameon=False, fontsize=8, + ) + figure.text( + .815, .41, + "Purple branches are privileged MPC teacher targets.\n" + "They are control-bounded but are not SOCP labels,\n" + "not certificates, and never enter ordinary D+.\n\n" + f"round: {shard.round_i}\n" + f"ordinary D: {len(shard.D)}\n" + f"ordinary D+: {len(shard.Dplus)}\n" + f"ordinary D-: {len(shard.Dminus)}\n" + f"teacher rows: {provenance['teacher_records']}\n" + f"teacher contexts: {provenance['teacher_contexts']}\n" + "teacher families: " + + ", ".join( + f"{name}={count}" + for name, count in provenance["teacher_families"].items() + ), + ha="left", va="top", fontsize=8, + ) + + def update(step): + for row, scenario in enumerate(scenarios): + for column, gamma in enumerate(gammas): + axis = axes[row, column] + axis.clear() + draw_cell( + axis, + lineages[(scenario, round(gamma, 8))], + teachers, + int(step), + ) + return [] + + for path in (output_png, output_json, output_mp4): + if path: + os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True) + if output_mp4: + movie = animation.FuncAnimation( + figure, update, frames=frames, + interval=1000 / int(fps), blit=False, + ) + movie.save( + output_mp4, + writer=animation.FFMpegWriter(fps=int(fps), bitrate=4600), + dpi=105, + ) + update(maximum) + figure.savefig(output_png, dpi=165, bbox_inches="tight") + plt.close(figure) + + report = dict( + status="SFM_B1_TEACHER_D_BRANCH_VIZ_COMPLETE", + round=int(shard.round_i), + scenarios=list(scenarios), + gammas=list(gammas), + ordinary_counts=dict( + D=len(shard.D), Dplus=len(shard.Dplus), Dminus=len(shard.Dminus), + ), + teacher_counts=dict( + records=provenance["teacher_records"], + contexts=provenance["teacher_contexts"], + families=provenance["teacher_families"], + ), + teacher_semantics=( + "control-bounded privileged Codex SFM-MPC CFM targets; separate " + "from ordinary verifier-labeled D/D+/D-; no safety label implied" + ), + context_match=( + "round shard SHA-256 plus exact context_id, scenario, gamma, " + "step, state, ped_xy, and ped_vel" + ), + scenario_selection=( + "explicit CLI scenarios" if scenarios_were_explicit + else "top three complete scenarios by teacher-row count" + ), + provenance=provenance, + source_manifest=teacher_bundle.get("provenance"), + frames=frames if output_mp4 else [maximum], + png=os.path.abspath(output_png), + mp4=None if not output_mp4 else os.path.abspath(output_mp4), + ) + _write_json(output_json, report) + return report + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--round-shard", required=True) + parser.add_argument("--teacher-buffer", required=True) + parser.add_argument("--output-png", required=True) + parser.add_argument("--output-json", required=True) + parser.add_argument("--output-mp4") + parser.add_argument("--scenarios", nargs=3, type=int) + parser.add_argument("--fps", type=int, default=5) + parser.add_argument("--frame-stride", type=int, default=2) + args = parser.parse_args(argv) + render( + args.round_shard, args.teacher_buffer, + args.output_png, args.output_json, + scenarios=args.scenarios, output_mp4=args.output_mp4, + fps=args.fps, frame_stride=args.frame_stride, + ) + + +if __name__ == "__main__": + main() diff --git a/overnight_run_07_12_sfm/sfm_kazuki.py b/overnight_run_07_12_sfm/sfm_kazuki.py index 424e03a..850d56d 100644 --- a/overnight_run_07_12_sfm/sfm_kazuki.py +++ b/overnight_run_07_12_sfm/sfm_kazuki.py @@ -797,6 +797,8 @@ def guided_generate(policy, ctx, state, goal, ped_pred, ped_vel, r_col, z_init, N, H = len(z_init), policy.H_pred z = z_init; ctxN = policy._expand_ctx(ctx, N) z_unguided = z_init.detach().clone() if collect_diagnostics else None + goal_integral = torch.zeros_like(z) if collect_diagnostics else None + safety_integral = torch.zeros_like(z) if collect_diagnostics else None safe_coef = _sample_safe_coefficients(cfg, z.device, z.dtype) markup = float(cfg.markup) ** torch.arange(H - 1, -1, -1, dtype=z.dtype, device=z.device) markup = markup[None, :, None] @@ -820,11 +822,18 @@ def guided_generate(policy, ctx, state, goal, ped_pred, ped_vel, r_col, z_init, raw_goal_norm = torch.linalg.vector_norm(g_goal) g_cbf = g_cbf * base_norm / (raw_cbf_norm + 1e-8) g_goal = g_goal * base_norm / (raw_goal_norm + 1e-8) - guidance = (float(cfg.goal_coef) * g_goal.reshape(N, H, 2) - + safe_coef * g_cbf.reshape(N, H, 2) * markup) + goal_guidance = float(cfg.goal_coef) * g_goal.reshape(N, H, 2) + safety_guidance = safe_coef * g_cbf.reshape(N, H, 2) * markup + guidance = goal_guidance + safety_guidance z = z + (tau_next - tau) * (base + guidance.reshape(N, -1)) if collect_diagnostics: z_unguided = z_unguided + (tau_next - tau) * base_unguided + goal_integral = goal_integral + ( + tau_next - tau + ) * goal_guidance.reshape(N, -1) + safety_integral = safety_integral + ( + tau_next - tau + ) * safety_guidance.reshape(N, -1) trace.append(dict( tau=tau, dtau=tau_next - tau, base_norm=float(base_norm.detach().cpu()), raw_cbf_grad_norm=float(raw_cbf_norm.detach().cpu()), @@ -834,6 +843,15 @@ def guided_generate(policy, ctx, state, goal, ped_pred, ped_vel, r_col, z_init, endpoint_min_pred_clear=float((torch.linalg.vector_norm( pos.unsqueeze(2) - ped_pred.unsqueeze(0), dim=3) - r_col).min().detach().cpu()), )) + if collect_diagnostics: + trace[-1].update( + integrated_goal_guidance=goal_integral.detach(), + integrated_safety_guidance=safety_integral.detach(), + component_semantics=( + "ODE-time integrals of the additive normalized goal and CBF " + "guidance fields evaluated along the guided latent path" + ), + ) return z, trace, z_unguided @@ -1096,13 +1114,37 @@ def kazuki_sfm_deploy(policy, episode, gamma, cfg=None, n_ped=20, T=180, reach=0 guided_final_action = U_gen[seed_index, 0].detach().cpu().numpy().astype(np.float32) unguided_final_action = U_unguided[seed_index, 0].detach().cpu().numpy().astype(np.float32) net_guidance_action = guided_final_action - unguided_final_action + component_row = ode_diag[-1] + integrated_goal = component_row.pop("integrated_goal_guidance") + integrated_safety = component_row.pop("integrated_safety_guidance") + component_semantics = component_row.pop("component_semantics") + goal_guidance_action = ( + integrated_goal[seed_index] + .reshape(H, 2)[0].detach().cpu().numpy().astype(np.float32) + * float(policy.u_max) + ) + safety_guidance_action = ( + integrated_safety[seed_index] + .reshape(H, 2)[0].detach().cpu().numpy().astype(np.float32) + * float(policy.u_max) + ) guidance_diag = dict( selected_generated_index=seed_index, guided_final_action=guided_final_action, unguided_final_action=unguided_final_action, net_guidance_action=net_guidance_action, net_guidance_norm=float(np.linalg.norm(net_guidance_action)), - semantics="guided final acceleration minus unguided final acceleration from the same latent sample") + goal_guidance_action=goal_guidance_action, + goal_guidance_norm=float(np.linalg.norm(goal_guidance_action)), + safety_guidance_action=safety_guidance_action, + safety_guidance_norm=float(np.linalg.norm(safety_guidance_action)), + component_semantics=component_semantics, + semantics=( + "net: guided final acceleration minus unguided final " + "acceleration from the same latent sample; goal/safety: " + "separate integrated additive guidance components before " + "MPPI refinement" + )) viz_controls = U_best.detach().clone() viz_controls[0] = torch.as_tensor(action, dtype=viz_controls.dtype, device=viz_controls.device) with torch.no_grad(): diff --git a/overnight_run_07_12_sfm/sfm_metrics2.py b/overnight_run_07_12_sfm/sfm_metrics2.py index c3762ed..5efc946 100644 --- a/overnight_run_07_12_sfm/sfm_metrics2.py +++ b/overnight_run_07_12_sfm/sfm_metrics2.py @@ -163,14 +163,15 @@ def certify_moving_window(segment, pedestrians, gamma, *, K=ARTIFICIAL_FACES, robot_c = robot - center radius = max(float(SS.R_SENSE), float(r_pad) * float(np.linalg.norm(robot_c, axis=1).max())) if len(robot) == 1: - # Reaching the goal at the current state creates an absorbing empty - # verification horizon. It is valid but never replay-eligible. + # This generic helper also audits the zero-transition tail of an + # already executed trajectory. Queried B1 plans never take this path: + # verify_query below requires and certifies all H=10 transitions. return True, [], dict( solver="exact_2d_angular_interval_socp", angular_grid=False, slack=float("inf"), worst_t=0, R_eff=float(radius), n_real=0, n_real_feasible=0, n_artificial=0, n_artificial_feasible=0, K_artificial=ARTIFICIAL_FACES, - empty_terminal_prefix=True, + empty_executed_tail=True, ) alpha = (1.0 - float(gamma)) ** np.arange(len(robot), dtype=float) beta = 1.0 - alpha @@ -201,34 +202,57 @@ def certify_moving_window(segment, pedestrians, gamma, *, K=ARTIFICIAL_FACES, ) -def verify_query(state, controls, ped_xy, ped_vel, gamma, *, reach=0.5): - """Resolve y without performance/cost terms; errors are explicit and non-storable.""" +def _verify_window(state, controls, ped_xy, ped_vel, gamma): + robot = rollout_positions(state, controls) + pedestrian = predict_pedestrians(ped_xy, ped_vel, H=len(controls)) + task = taskspace_ok(robot) + collision = collision_free_time_indexed(robot, pedestrian) + certificate, faces, diagnostics = certify_moving_window( + robot, pedestrian, gamma, + ) + y = bool(task and collision and certificate) + return dict( + resolved=True, error=None, y=int(y), taskspace=bool(task), + collision_free=bool(collision), certificate=bool(certificate), + segment=robot, pedestrian_prediction=pedestrian, faces=faces, + diagnostics=diagnostics, + ) + + +def verify_executed_window(state, controls, ped_xy, ped_vel, gamma): + """Certify one terminal-truncated executed window of length 1 through 10. + + This API is for offline trajectory evaluation. A short window is complete + relative to the executed trajectory; it is not a B1 terminal-prefix query + and therefore deliberately has no ``full_h`` or ``train_eligible`` field. + """ + try: + controls = np.asarray(controls, np.float32).reshape(-1, 2) + if not 1 <= len(controls) <= 10: + raise ValueError("executed-window verifier requires 1 <= H_t <= 10") + result = _verify_window(state, controls, ped_xy, ped_vel, gamma) + result.update(window_horizon=len(controls)) + return result + except Exception as error: + return dict(resolved=False, error=f"{type(error).__name__}: {error}") + + +def verify_query(state, controls, ped_xy, ped_vel, gamma): + """Certify every queried plan over all H=10 transitions. + + Goal reach is a closed-loop episode trigger after the selected first action; + it never truncates a candidate window or changes its verifier label. + """ try: controls = np.asarray(controls, np.float32).reshape(-1, 2) if len(controls) != 10: raise ValueError("B1 verifier requires H=10") - robot = rollout_positions(state, controls) - pedestrian = predict_pedestrians(ped_xy, ped_vel, H=len(controls)) - goal_distance = np.linalg.norm(robot - SS.GOAL[None], axis=1) - reached = np.flatnonzero(goal_distance < float(reach)) - terminal_step = int(reached[0]) if len(reached) else len(controls) - # A goal hit defines an absorbing terminal prefix. Post-goal repeats are not verified or replayed. - prefix_robot = robot[:terminal_step + 1] - prefix_pedestrian = pedestrian[:terminal_step + 1] - task = taskspace_ok(prefix_robot) - collision = collision_free_time_indexed(prefix_robot, prefix_pedestrian) - certificate, faces, diagnostics = certify_moving_window( - prefix_robot, prefix_pedestrian, gamma, - ) - y = bool(task and collision and certificate) - return dict( - resolved=True, error=None, y=int(y), taskspace=bool(task), - collision_free=bool(collision), certificate=bool(certificate), - full_h=bool(terminal_step == len(controls)), terminal_step=terminal_step, - train_eligible=bool(y and terminal_step == len(controls)), - segment=robot, pedestrian_prediction=pedestrian, faces=faces, - diagnostics=diagnostics, + result = _verify_window(state, controls, ped_xy, ped_vel, gamma) + result.update( + full_h=True, terminal_step=len(controls), + train_eligible=bool(result["y"]), ) + return result except Exception as error: # worker boundary: a failed solver/query enters no store. return dict(resolved=False, error=f"{type(error).__name__}: {error}") diff --git a/overnight_run_07_12_sfm/show_claude_neutral_handoff.py b/overnight_run_07_12_sfm/show_claude_neutral_handoff.py new file mode 100644 index 0000000..30fc8dd --- /dev/null +++ b/overnight_run_07_12_sfm/show_claude_neutral_handoff.py @@ -0,0 +1,95 @@ +"""Print the frozen Claude handoff from existing Helios artifacts only. + +This command never launches a rollout or recomputes a metric. It authenticates +the shared visualization, checkpoint, and fresh disjoint-M50 delivery before +printing the two requested reference rows. +""" +from __future__ import annotations + +import hashlib +import json +from pathlib import Path + + +VIDEO = Path( + "/data3/research1/sfm_neutral_claude_handoff_assets/" + "neutral_continuation.mp4" +) +VIDEO_SHA256 = "c9437021d5a2c25f2c735521f00ded0727b819f4ab2beca3e97c4beb5370b45f" +DELIVERY = Path( + "/data3/research1/sfm_neutral_gamma_temp_0e441d6/" + "DELIVERY_COMPLETE.json" +) +DELIVERY_STATUS = "SFM_NEUTRAL_GAMMA_TEMPERATURE_M50_COMPLETE" +METRICS_SOURCE_COMMIT = "0e441d68644e89017a4c312ad59e7b520ca80208" +CHECKPOINT = Path( + "/home/dohyun/projects/sfm_hp10_b1_runs/103476d/pretrained_hp10.pt" +) +CHECKPOINT_SHA256 = "1b5179c935d3eeff8824967d707d64cc9bab273949ee1f0e4f190172bab1b215" + + +def sha256_file(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as stream: + for chunk in iter(lambda: stream.read(1 << 20), b""): + digest.update(chunk) + return digest.hexdigest() + + +def require_file(path: Path, expected_sha256: str | None = None) -> None: + if not path.is_file(): + raise FileNotFoundError(path) + if expected_sha256 is not None: + observed = sha256_file(path) + if observed != expected_sha256: + raise RuntimeError( + f"SHA-256 mismatch for {path}: {observed} != {expected_sha256}" + ) + + +def main() -> None: + require_file(VIDEO, VIDEO_SHA256) + require_file(CHECKPOINT, CHECKPOINT_SHA256) + require_file(DELIVERY) + payload = json.loads(DELIVERY.read_text()) + if payload.get("status") != DELIVERY_STATUS: + raise RuntimeError("cached M50 delivery is incomplete") + if payload.get("source_commit") != METRICS_SOURCE_COMMIT: + raise RuntimeError("cached M50 delivery source changed") + bank = payload["fresh_confirmation_bank"] + if (bank["M_per_gamma"], bank["ep0"]) != (50, 485000): + raise RuntimeError("fresh disjoint M50 bank changed") + by_method = {row["method"]: row for row in payload["final_records"]} + if set(("pretrained", "kazuki_locked")) - set(by_method): + raise RuntimeError("cached reference methods are missing") + + print(f"NEUTRAL_CONTINUATION_MP4={VIDEO}") + print(f"NEUTRAL_CONTINUATION_MP4_SHA256={VIDEO_SHA256}") + print(f"PRETRAINED_CHECKPOINT={CHECKPOINT}") + print(f"PRETRAINED_CHECKPOINT_SHA256={CHECKPOINT_SHA256}") + print( + "CACHED_REFERENCE_BANK=" + f"fresh disjoint M={bank['M_per_gamma']}/gamma, ep0={bank['ep0']}, " + f"noise_seed={bank['noise_seed']}" + ) + print("method\tSR\tCR\ttimeout\tValidity\tclearance_m\ttime_to_goal_s") + for method in ("pretrained", "kazuki_locked"): + row = by_method[method] + values = row["pooled"] + print( + f"{method}\t{values['SR']:.6f}\t{values['CR']:.6f}\t" + f"{values['timeout']:.6f}\t{values['Validity']:.6f}\t" + f"{values['clearance']:.6f}\t{values['time_to_goal']:.6f}" + ) + schedule = row.get("temperature_by_gamma") + if schedule is not None: + print( + f"{method}_temperature_by_gamma=" + f"{dict(zip(('0.1','0.2','0.3','0.4','0.5','0.7','1.0'), schedule))}" + ) + print(f"CACHED_REFERENCE_DELIVERY={DELIVERY}") + print("No rollout or metric was recomputed.") + + +if __name__ == "__main__": + main() diff --git a/overnight_run_today/src/flow_policy.py b/overnight_run_today/src/flow_policy.py index 40110bd..39066ea 100644 --- a/overnight_run_today/src/flow_policy.py +++ b/overnight_run_today/src/flow_policy.py @@ -97,6 +97,21 @@ def sample(self, n: int, ctx: torch.Tensor, nfe: int = 12, U = (x.reshape(n, self.T, 2) * self.u_max).clamp(-self.u_max, self.u_max) return U + @torch.no_grad() + def phi_s_from_x0( + self, U_controls: torch.Tensor, ctx: torch.Tensor, + x0: torch.Tensor, s: float = 0.9, + ) -> torch.Tensor: + """Noised-flow representation using each proposal's original base noise.""" + B = U_controls.shape[0] + x1 = (U_controls / self.u_max).reshape(B, self.d) + if tuple(x0.shape) != (B, self.d): + raise ValueError(f"x0 shape {tuple(x0.shape)} != {(B, self.d)}") + x0 = x0.to(device=x1.device, dtype=x1.dtype) + x_s = (1 - float(s)) * x0 + float(s) * x1 + tau = torch.full((B,), float(s), device=x1.device) + return self.features(x_s, tau, self._expand_ctx(ctx, B)) + @torch.no_grad() def phi_s(self, U_controls: torch.Tensor, ctx: torch.Tensor, s: float = 0.9) -> torch.Tensor: """Noised-flow representation at level s, averaged over fixed noise templates -> [B, width]."""