Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
74 commits
Select commit Hold shift + click to select a range
bf53dee
Restart SFM audit with full-H continuation labels
DHLeexpress Jul 23, 2026
830f5a8
Record authenticated full-episode audit
DHLeexpress Jul 23, 2026
58ec896
Add two-round SFM alpha replay sweep
DHLeexpress Jul 23, 2026
2522c2c
Aggregate two-round SFM factorial results
DHLeexpress Jul 23, 2026
041d982
Add offline executed-window SFM expansion sweep
DHLeexpress Jul 24, 2026
9a6595e
Queue offline SFM sweep after exclusive GPU smoke
DHLeexpress Jul 24, 2026
3abd258
Enforce strict gamma-balanced GP support
DHLeexpress Jul 24, 2026
4c8b41d
Preserve proposal base noise in SFM uncertainty
DHLeexpress Jul 24, 2026
f36393b
Fix signed replay batching and target GPUs 1 and 3
DHLeexpress Jul 24, 2026
c63ac43
Add x0-faithful partial branch comparison
DHLeexpress Jul 24, 2026
db80dfc
Add native SafeMPPI-cost offline selector sweep
DHLeexpress Jul 24, 2026
9ed6b4c
Add safety-performance rank selector diagnostic
DHLeexpress Jul 24, 2026
5fb90b0
Allow balanced selector in shared B1 config
DHLeexpress Jul 24, 2026
97630c6
Extend paired comparison to balanced selector
DHLeexpress Jul 24, 2026
3f80746
Recover completed SFM offline evaluations
DHLeexpress Jul 25, 2026
5811bb6
Accept authenticated legacy margin recipes in recovery
DHLeexpress Jul 25, 2026
1e5ee41
Allow single-GPU offline evaluation recovery
DHLeexpress Jul 25, 2026
f8d098e
Add staged offline evaluation funnel
DHLeexpress Jul 25, 2026
f06e8dd
Automate disjoint final M100 confirmation
DHLeexpress Jul 26, 2026
779d2de
Claude recipe study: opt-in replay interventions (hard A/B population…
DHLeexpress Jul 26, 2026
2f4b68e
Stage A results + orig_plus_recovery mode (declared pre-evaluation): …
DHLeexpress Jul 26, 2026
1b161fe
Freeze Stage-C recipe (B2 cost/hard by predeclared rule) + SELECTION_…
DHLeexpress Jul 26, 2026
511c120
Recovery family v2 (goal-directed dodge-then-cruise, declared pre-eva…
DHLeexpress Jul 26, 2026
06abdca
Stage D complete: frozen cost/hard recipe = honest null (rule-selecte…
DHLeexpress Jul 26, 2026
9d2bc9e
Iteration-2 freeze applied per declared criterion: R_A (B9 knobs, v1 …
DHLeexpress Jul 26, 2026
0b40054
Iteration-2 M50 selection: r1 only eligible round, dCR -.123 dV +.141…
DHLeexpress Jul 26, 2026
faceb80
MPC distillation study scaffolding: privileged Codex pool import (rea…
DHLeexpress Jul 26, 2026
0c62905
MPC study pre-registration: arms, banks (M10 350000 audit+select, M50…
DHLeexpress Jul 26, 2026
d42e913
Recipe study DELIVERY: M100 confirmation all-CI-clean (dCR -.100, dV …
DHLeexpress Jul 26, 2026
123c2b9
MPC distillation study COMPLETE (fail-closed at M10): visible raw-pol…
DHLeexpress Jul 27, 2026
455fd58
MPC pool: allow autograd through Codex guidance (deploy-exact ctx squ…
DHLeexpress Jul 27, 2026
4b84bdb
R1 stable-continuation study: pre-registration (banks 380k anchor / 3…
DHLeexpress Jul 28, 2026
c641c95
continuation engine: create per-round checkpoint directory before sav…
DHLeexpress Jul 28, 2026
78b4422
R1 continuation: development stopped per pre-registered rule (0/16 ad…
DHLeexpress Jul 28, 2026
27d40a6
R1 continuation FINAL: no admissible continuation (timeout gate binds…
DHLeexpress Jul 28, 2026
74fe199
Add isolated unverified MPC teacher ablation
DHLeexpress Jul 28, 2026
2ed4373
Fix GRU conflict probe training mode
DHLeexpress Jul 28, 2026
2a87332
Use canonical D-branch label colors
DHLeexpress Jul 28, 2026
043f1b3
Teacher-continuation study: pre-registration (gates/doses/banks 420k/…
DHLeexpress Jul 28, 2026
40bbd5d
reuse-buffer: keep saved audit byte-identical (provenance in COMPLETE…
DHLeexpress Jul 28, 2026
ec41466
teacher driver: assert canonical CONTENT hash of D_MPC records (torch…
DHLeexpress Jul 28, 2026
ce4fa9f
Teacher study COMPLETE fail-closed at round 1: ordinary e10 step alon…
DHLeexpress Jul 28, 2026
736fa6c
Corrected privileged-distillation study: pre-registration (banks 440k…
DHLeexpress Jul 29, 2026
0c52010
AFE-fidelity controller benchmark: pre-registration (banks 460k/390k/…
DHLeexpress Jul 29, 2026
5ce8a03
AFE expansion+funnel stage: 2 promoted recipes x 10 rounds (W=2 certi…
DHLeexpress Jul 29, 2026
133ce7d
AFE-fidelity funnel FINAL: phase-0 support paradox (per-context suppo…
DHLeexpress Jul 29, 2026
1e405a0
UNFAITHFUL pilot (user-directed): unverified goal-directed controller…
DHLeexpress Jul 30, 2026
a43a8bd
UNFAITHFUL full launch p05/E4/lr1e-5 x10 rounds, resumable, M10 every…
DHLeexpress Jul 30, 2026
27c6b63
Add same-latent Kazuki repair acquisition audit
DHLeexpress Jul 30, 2026
9e42797
Fix six-row provenance source lookup
DHLeexpress Jul 30, 2026
a7be9ed
Add neutral continuation for negative guided repairs
DHLeexpress Jul 31, 2026
c0b90a4
Fix neutral terminal labels in six-row render
DHLeexpress Jul 31, 2026
0e86fac
Add isolated neutral teacher round-one sanity study
DHLeexpress Jul 31, 2026
9774ad0
Harden neutral teacher audit semantics
DHLeexpress Jul 31, 2026
12a7559
Fix train mode for neutral gradient audit
DHLeexpress Jul 31, 2026
a36dfe7
Add paired neutral multiround correction study
DHLeexpress Jul 31, 2026
52ba2c7
Add locked-temperature disjoint M50 evaluation
DHLeexpress Jul 31, 2026
728de59
Fail closed on invalid temperature candidates
DHLeexpress Jul 31, 2026
0a13ca9
Add gamma calibration and exact neutral continuation
DHLeexpress Jul 31, 2026
b59bf6e
Automate exact r100 continuation after failed calibration
DHLeexpress Jul 31, 2026
9c0ec8a
Gate gamma temperature schedules on trend
DHLeexpress Jul 31, 2026
9240602
Require final liveness and gamma trend gates
DHLeexpress Jul 31, 2026
cd40a71
Add fail-closed neutral creative sanity stages
DHLeexpress Jul 31, 2026
fce7576
Harden autonomous neutral creative follow-up
DHLeexpress Jul 31, 2026
8cd65be
Restore global-temperature evaluator parity
DHLeexpress Jul 31, 2026
95f354a
Fail closed on temperature parity drift
DHLeexpress Jul 31, 2026
6838482
Reuse locked calibration baseline cells
DHLeexpress Jul 31, 2026
57ce79b
Harden autonomous neutral expansion follow-up
DHLeexpress Jul 31, 2026
1f26b53
Authenticate expansion evaluation lineage
DHLeexpress Jul 31, 2026
0e441d6
Support legacy scalar calibration metadata
DHLeexpress Jul 31, 2026
73a9c0a
Add ESS acquisition support diagnostic
DHLeexpress Jul 31, 2026
6c99a05
Add authenticated Claude neutral expansion handoff
DHLeexpress Jul 31, 2026
fd2a719
Permit diagnostic ESS targets in repair collector
DHLeexpress Jul 31, 2026
dd98b33
Record ESS acquisition provenance in plot
DHLeexpress Jul 31, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
141 changes: 141 additions & 0 deletions FRESH_RESTART.md
Original file line number Diff line number Diff line change
@@ -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`.
138 changes: 138 additions & 0 deletions overnight_run_07_12_sfm/CLAUDE_NEUTRAL_EXPANSION_HANDOFF.md
Original file line number Diff line number Diff line change
@@ -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/<CLAUDE_PRIVATE_WORKTREE>
/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_<frozen_sha7>/`. 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.
112 changes: 112 additions & 0 deletions overnight_run_07_12_sfm/analysis/test_claude_afe_fidelity.py
Original file line number Diff line number Diff line change
@@ -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
Loading