"""Validation-only selection: score every p2_long export, uniform soups of its last exports, the p1_short final export and p1-final/p2-final weight averages on the 7,824-row selection split; freeze the lowest-NLL candidate. The audit partition and the benchmark are not touched here.""" import gc, json, shutil, sys import numpy as np import torch from datasets import load_from_disk from safetensors.torch import load_file, save_file from common import * import engine from train import exports SEL = CKPT/'selection' def soup(paths, dest, weights=None): """Weighted (default uniform) average of BF16 exports, computed in FP32, saved as BF16.""" dest = Path(dest) if (dest/'model.safetensors').exists(): return dest weights = weights or [1/len(paths)]*len(paths) assert abs(sum(weights) - 1) < 1e-6 acc = None for p, w in zip(paths, weights): state = load_file(str(Path(p)/'model.safetensors')) if acc is None: acc = {k: v.float()*w for k, v in state.items()} else: assert acc.keys() == state.keys() for k, v in state.items(): acc[k] += v.float()*w stage = dest.with_name(dest.name+'.staging') if stage.exists(): shutil.rmtree(stage) shutil.copytree(Path(paths[0]), stage, ignore=shutil.ignore_patterns('model.safetensors', '*.npz', 'training_progress.json', 'soup.json')) save_file({k: v.to(torch.bfloat16).contiguous() for k, v in acc.items()}, str(stage/'model.safetensors'), metadata={'format': 'pt'}) atomic_json(stage/'soup.json', {'members': [str(p) for p in paths], 'weights': weights}) stage.rename(dest) return dest def evaluate(path, name, val, indices): out = SEL/(name+'.npz') model = engine.load_model(path) logits = engine.predict(model, val, indices, 'selection-'+name, out) del model; gc.collect(); torch.cuda.empty_cache() r = engine.report(val, indices, logits) return {'name': name, 'path': str(path), 'sha256': sha(Path(path)/'model.safetensors'), 'nll': r['overall']['nll'], 'balanced_accuracy': r['overall']['balanced_accuracy'], 'accuracy': r['overall']['accuracy'], 'false_approve_count': r['overall']['false_approve_count'], 'false_deny_count': r['overall']['false_deny_count'], 'long_accuracy': r['slices']['length']['16k-64k']['accuracy'], 'auroc': r['overall']['auroc'], 'report': r} def main(run='p2_long'): if (SEL/'summary.json').exists() and '--force' not in sys.argv: print(json.dumps(json.loads((SEL/'summary.json').read_text())['frozen'], indent=1)); return SEL.mkdir(parents=True, exist_ok=True) val = load_from_disk(str(DATA/'validation')); indices = np.load(DATA/'validation_partitions.npz')['selection'] assert (CKPT/run/'complete.json').exists() ex = exports(run); results = {} for p in ex: results[f'{run}-{p.name}'] = evaluate(p, f'{run}-{p.name}', val, indices) p1 = exports('p1_short')[-1] results[f'p1_short-{p1.name}'] = evaluate(p1, f'p1_short-{p1.name}', val, indices) for k in (2, 3, 4): if len(ex) >= k: name = f'soup-{run}-last{k}' results[name] = evaluate(soup(ex[-k:], SEL/name), name, val, indices) # Added before any benchmark evaluation (p2 raised validation NLL on short inputs): p1-final/p2-final averages. for w in (0.3, 0.5, 0.7): name = f'soup-p1final{w:.1f}+p2final{1-w:.1f}' results[name] = evaluate(soup([p1, ex[-1]], SEL/name, [w, round(1-w, 1)]), name, val, indices) eligible = list(results) ranked = sorted(eligible, key=lambda n: results[n]['nll']) frozen = results[ranked[0]] for r in sorted(results.values(), key=lambda r: r['nll']): print(f"{'*' if r['name'] == frozen['name'] else ' '}{r['name']:32s} nll={r['nll']:.5f} bacc={r['balanced_accuracy']:.5f} acc={r['accuracy']:.5f} " f"FA={r['false_approve_count']:3d} FD={r['false_deny_count']:3d} long={r['long_accuracy']:.4f}", flush=True) atomic_json(SEL/'summary.json', {'frozen': {k: v for k, v in frozen.items() if k != 'report'}, 'ranked': ranked, 'criterion': 'lowest NLL on the 7,824-row validation selection split among p2_long exports, last-2/3/4 uniform p2 soups, the p1_short final export and p1-final/p2-final weight averages (0.3/0.5/0.7); frozen before any benchmark or audit evaluation', 'candidates': [{k: v for k, v in r.items() if k != 'report'} for r in sorted(results.values(), key=lambda r: r['nll'])]}) for r in results.values(): atomic_json(SEL/(r['name']+'.report.json'), r['report']) event('frozen', **{k: v for k, v in frozen.items() if k != 'report'}) if __name__ == '__main__': main()