Leplanner / code /report /make_tables.py
nottygian's picture
Push code package
872cf4d verified
Raw
History Blame Contribute Delete
11 kB
r"""Export every table in the report as LaTeX, straight from the raw results.
Same contract as ``make_figures.py``: read ``data/runs/eval/results.jsonl``,
``data/runs/diagnostics/*`` and the checkpoints, emit ``report/tables/*.tex``
for ``\input``. Nothing in the document is hand-typed, so the prose cannot
drift from the numbers.
python report/make_tables.py
"""
import json
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
OUT = Path(__file__).resolve().parent / 'tables'
OUT.mkdir(parents=True, exist_ok=True)
sys.path.insert(0, str(Path(__file__).resolve().parent))
from make_figures import load_eval, load_diag, latest # noqa: E402
TEX = {
'abl_terminal_only': r'terminal-only ($\alpha=0$)',
'original': 'original',
'abl_no_support': r'no support ($\lambda_{\mathrm{sup}}=0$)',
'ah_hold0.0': r'arrival ($\lambda_h=0$)',
'ah_hold0.5': r'\textbf{arrival+hold} ($\lambda_h=0.5$)',
'ah_hold1.0': r'arrival+hold ($\lambda_h=1$)',
'cem': r'CEM ($300\times30$)',
}
ORDER = ['abl_terminal_only', 'original', 'abl_no_support',
'ah_hold0.0', 'ah_hold1.0', 'ah_hold0.5', 'cem']
def write(name, body):
(OUT / f'{name}.tex').write_text(body)
print(f' wrote tables/{name}.tex')
def t_main(rows):
"""Headline matrix: success at both schedules, the gap, and cost."""
lines = [
r'\begin{tabular}{lrrrrr}',
r'\toprule',
r'variant & $m{=}1$ & $m{=}5$ & gap & rows/ep & rows/call \\',
r'\midrule',
]
for v in ORDER:
K = None if v == 'cem' else 3
a, b = latest(rows, v, 1, K=K), latest(rows, v, 5, K=K)
if not (a and b):
continue
gap = a['success_rate'] - b['success_rate']
pc = a['predictor_rows_per_episode'] * a['num_eval'] / a['predictor_calls']
mark = r'\phantom{-}' if gap >= 0 else ''
lines.append(
f"{TEX[v]} & {a['success_rate']:.0f} & {b['success_rate']:.0f} & "
f"{mark}{gap:+.0f} & {a['predictor_rows_per_episode']:.1f} & "
f"{pc:.1f} \\\\"
)
if v == 'ah_hold0.5':
lines.append(r'\midrule')
lines += [r'\bottomrule', r'\end{tabular}']
write('main_results', '\n'.join(lines))
def t_contraction(diag):
lines = [
r'\begin{tabular}{lrrrr}',
r'\toprule',
r'variant & $c$ & $b$ & $D^\ast$ & $R^2$ \\',
r'\midrule',
]
for v in ORDER[:-1]:
d = diag.get(v)
if not d:
continue
f = d['contraction']['exec1']
lines.append(
f"{TEX[v]} & {f['c']:.4f} & {f['b']:.4f} & "
f"{f['fixed_point']:.4f} & {f['r2']:.3f} \\\\"
)
lines += [r'\bottomrule', r'\end{tabular}']
write('contraction', '\n'.join(lines))
def t_paired(path=None):
"""Paired comparisons, recomputed with the same tests used in the run.
Every eval row shares the same 50 seeded held-out episodes, so the
comparison is paired: exact McNemar on the discordant episodes, plus a
percentile bootstrap CI on the difference.
"""
import numpy as np
sys.path.insert(0, str(ROOT / 'scripts'))
from paired_stats import mcnemar_exact, bootstrap_ci, label # noqa: E402
rows = [json.loads(x) for x in
(ROOT / 'data/runs/eval/results.jsonl').read_text().splitlines()
if x.strip()]
by = {label(r): np.array(r['episode_successes'], dtype=bool)
for r in rows if r.get('episode_successes')}
# (section heading, human-readable question, row A, row B)
comps = [
('The pathology', r'original: $m{=}1$ vs $m{=}5$',
'controller_K3', 'controller_K3+exec5'),
(None, r'terminal-only: $m{=}1$ vs $m{=}5$',
'controller_K3[terminal_only]', 'controller_K3+exec5[terminal_only]'),
(None, r'no support: $m{=}1$ vs $m{=}5$',
'controller_K3[no_support]', 'controller_K3+exec5[no_support]'),
(None, r'CEM: $m{=}1$ vs $m{=}5$',
'cem_s300_n30', 'cem_s300_n30+exec5'),
('The fix removes it', r'arrival+hold $\lambda_h{=}0.5$: $m{=}1$ vs $m{=}5$',
'controller_K3[ah_hold0.5]', 'controller_K3+exec5[ah_hold0.5]'),
(None, r'arrival $\lambda_h{=}0$: $m{=}1$ vs $m{=}5$',
'controller_K3[ah_hold0.0]', 'controller_K3+exec5[ah_hold0.0]'),
(None, r'arrival+hold $\lambda_h{=}1$: $m{=}1$ vs $m{=}5$',
'controller_K3[ah_hold1.0]', 'controller_K3+exec5[ah_hold1.0]'),
('Ablations at $m{=}1$', r'arrival+hold $0.5$ vs original',
'controller_K3[ah_hold0.5]', 'controller_K3'),
(None, r'arrival+hold $0.5$ vs terminal-only',
'controller_K3[ah_hold0.5]', 'controller_K3[terminal_only]'),
(None, r'terminal-only vs original',
'controller_K3[terminal_only]', 'controller_K3'),
(None, r'no support vs original',
'controller_K3[no_support]', 'controller_K3'),
('Hold weight', r'$\lambda_h{=}0.5$ vs $\lambda_h{=}0$',
'controller_K3[ah_hold0.5]', 'controller_K3[ah_hold0.0]'),
(None, r'$\lambda_h{=}0.5$ vs $\lambda_h{=}1$',
'controller_K3[ah_hold0.5]', 'controller_K3[ah_hold1.0]'),
(None, r'$\lambda_h{=}0$ vs $\lambda_h{=}1$',
'controller_K3[ah_hold0.0]', 'controller_K3[ah_hold1.0]'),
('Against the best baselines', r'arrival+hold $0.5$ vs CEM at $m{=}5$',
'controller_K3[ah_hold0.5]', 'cem_s300_n30+exec5'),
(None, r'arrival+hold $0.5$ vs original at $m{=}5$',
'controller_K3[ah_hold0.5]', 'controller_K3+exec5'),
(None, r'arrival+hold $0.5$ vs original $K{=}3$, $m{=}4$',
'controller_K3[ah_hold0.5]', 'controller_K3+exec4'),
]
lines = [
r'\begin{tabular}{lrrrc}',
r'\toprule',
r'comparison & $\Delta$ (pts) & 95\% CI & $p$ & \\',
r'\midrule',
]
dump = []
first = True
for section, name, ka, kb in comps:
if ka not in by or kb not in by:
print(f' !! skip {ka} vs {kb}: missing')
continue
if section:
if not first:
lines.append(r'\addlinespace')
lines.append(rf'\multicolumn{{5}}{{l}}{{\emph{{{section}}}}} \\')
first = False
sa, sb = by[ka], by[kb]
delta = (sa.mean() - sb.mean()) * 100
lo, hi = bootstrap_ci(sa, sb)
pval, _, _ = mcnemar_exact(sa, sb)
star = r'$\ast$' if pval < 0.05 else ''
pstr = r'$<10^{-4}$' if pval < 1e-4 else f'{pval:.4f}'
lines.append(
rf'\quad {name} & {delta:+.0f} & '
rf'$[{lo:+.0f},\,{hi:+.0f}]$ & {pstr} & {star} \\'
)
dump.append({'label': name, 'a': ka, 'b': kb, 'delta': float(delta),
'ci': [float(lo), float(hi)], 'p': float(pval)})
lines += [r'\bottomrule', r'\end{tabular}']
write('paired_stats', '\n'.join(lines))
# persist the numbers so the prose can be checked against them later
out = ROOT / 'data/runs/eval/paired_stats.jsonl'
out.write_text('\n'.join(json.dumps(d) for d in dump) + '\n')
print(f' wrote {out.relative_to(ROOT)}')
def t_support(rows=None):
p = ROOT / 'data/runs/diagnostics/viol_posthoc.json'
if not p.exists():
print(' !! viol_posthoc.json missing; skipping support table')
return
rec = json.loads(p.read_text())
key = {'controller': 'original'}
lines = [
r'\begin{tabular}{lrrr}',
r'\toprule',
r'variant & violation frac. & $\mathcal{L}_{\mathrm{sup}}$ & '
r'mean NLL/dim \\',
r'\midrule',
]
for name, d in rec.items():
v = key.get(name, name)
lines.append(
f"{TEX.get(v, v)} & {d['violation_fraction']:.3f} & "
f"{d['support_loss']:.4f} & {d['mean_nll_per_dim']:.3f} \\\\"
)
c95 = next(iter(rec.values()))['c95']
lines += [
r'\midrule',
rf'\multicolumn{{4}}{{l}}{{\footnotesize threshold '
rf'$c_{{95}}={c95:.4f}$ (95th percentile of demonstration NLL/dim)}} \\',
r'\bottomrule', r'\end{tabular}',
]
write('support', '\n'.join(lines))
def t_training(rows=None):
"""Final validation losses per checkpoint, read from the saved state.
The original controller predates the arrival logging, so its $d_q$ comes
from ``recover_profiles.py`` (same ``evaluate()``, same held-out split)
and is marked with a dagger.
"""
import torch
posthoc = {}
p = ROOT / 'data/runs/diagnostics/profiles_posthoc.json'
if p.exists():
posthoc = json.loads(p.read_text())
names = [('abl_terminal_only', 'abl_terminal_only'),
('controller', 'original'),
('abl_no_support', 'abl_no_support'),
('ah_hold0.0', 'ah_hold0.0'),
('ah_hold1.0', 'ah_hold1.0'),
('ah_hold0.5', 'ah_hold0.5')]
lines = [
r'\begin{tabular}{lrrrrr}',
r'\toprule',
r'variant & $\alpha$ & $\lambda_{\mathrm{sup}}$ & $\lambda_h$ & '
r'val $d_H$ & val $d_q$ \\',
r'\midrule',
]
dag = False
for dirname, v in names:
p = ROOT / f'data/runs/{dirname}/controller.pt'
if not p.exists():
continue
ck = torch.load(p, map_location='cpu', weights_only=False)
a = ck.get('args', {})
hw = a.get('hold_weight') if a.get('arrival_hold') else None
arr = ck.get('val_arrival')
arr_s = '--'
if arr is not None:
arr_s = f'{arr:.4f}'
elif dirname in posthoc:
arr_s = rf"{posthoc[dirname]['arrival']:.4f}$^\dagger$"
dag = True
lines.append(
f"{TEX[v]} & {a.get('alpha', 0):.2f} & "
f"{a.get('lambda_support', 0):.2f} & "
f"{'--' if hw is None else f'{hw:.1f}'} & "
f"{ck.get('val_terminal', float('nan')):.4f} & {arr_s} \\\\"
)
if dag:
lines += [
r'\midrule',
r'\multicolumn{6}{l}{\footnotesize $\dagger$ recomputed '
r'post-hoc; this run predates the arrival logging} \\',
]
lines += [r'\bottomrule', r'\end{tabular}']
write('training', '\n'.join(lines))
def main():
print('loading results...')
rows = load_eval()
diag = load_diag()
print(f' {len(rows)} eval rows, {len(diag)} diagnostics records')
print('writing tables...')
t_main(rows)
t_contraction(diag)
t_paired()
t_support()
try:
t_training()
except Exception as e:
print(f' !! could not read checkpoints ({e}); skipping training table')
print('done.')
if __name__ == '__main__':
sys.exit(main())