auto-200m-2 / training /choose.py
ProCreations's picture
auto-200m-2: ModernBERT-base approve/deny gate, 65,536-token context
0bbb929 verified
Raw History Blame Contribute Delete
4.69 kB
"""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()