auto-200m-2 / training /train.py
ProCreations's picture
auto-200m-2: ModernBERT-base approve/deny gate, 65,536-token context
0bbb929 verified
Raw History Blame Contribute Delete
4.11 kB
"""Run the LR pilots or one planned training phase. Usage: train.py pilots | train.py RUN_NAME"""
import json, sys
import numpy as np
from datasets import load_from_disk
from common import *
import engine
def exports(run):
folder = CKPT/run
return sorted([p for p in folder.glob('step-*') if (p/'model.safetensors').exists()], key=lambda p: int(p.name.split('-')[-1]))
def resolve_start(spec):
if spec == 'init': return init_path()
kind, run = spec.split(':')
assert kind == 'final' and (CKPT/run/'complete.json').exists(), f'{run} not complete'
return exports(run)[-1]
def index_sets(run, seed):
ds = load_from_disk(str(DATA/'train_all'))
lengths = np.array(ds['length']); rows = np.array(ds['source_row'])
aug = rows >= 1_000_000
short_orig = np.flatnonzero(~aug & (lengths <= 4096)); long_orig = np.flatnonzero(~aug & (lengths > 4096)); aug_idx = np.flatnonzero(aug)
rng = np.random.default_rng(seed)
name = run['indices']
if name == 'short_all': return np.sort(np.concatenate([short_orig, aug_idx]))
if name == 'long_mix':
replay = rng.choice(short_orig, int(len(long_orig)*run['short_replay']), replace=False)
aug_replay = rng.choice(aug_idx, int(len(long_orig)*run['aug_replay']), replace=False)
return np.sort(np.concatenate([long_orig, replay, aug_replay]))
if name == 'pilot':
return np.sort(rng.choice(np.concatenate([short_orig, aug_idx]), run['rows'], replace=False))
raise ValueError(name)
def pilot_lr():
result = json.loads((CKPT/'pilot_choice.json').read_text())
return result['lr']
def resolve_lr(spec):
if isinstance(spec, (int, float)): return float(spec)
if spec == 'pilot': return pilot_lr()
base, factor = spec.split('*'); assert base == 'pilot'
return pilot_lr()*float(factor)
def run_phase(run):
start = resolve_start(run['start_from'])
config = {**run, 'lr_resolved': resolve_lr(run['lr']), 'start_path': str(start)}
indices = index_sets(run, SEED + 101)
return engine.train_phase(run['name'], start, indices, epochs=run['epochs'], lr=config['lr_resolved'],
teacher_logits=DATA/'teacher_logits_all.npy', alpha=run['alpha'], temperature=run['temperature'],
evals_per_epoch=run.get('evals_per_epoch', 4), seed_offset=run.get('seed_offset', 0),
deny_weight=run.get('deny_weight', 1.0), config=config, export=run.get('export', True))
def pilots():
if (CKPT/'pilot_choice.json').exists(): print(json.loads((CKPT/'pilot_choice.json').read_text())); return
p = plan()['pilot']; results = []
val = load_from_disk(str(DATA/'validation')); monitor = np.load(DATA/'validation_partitions.npz')['monitor']
lengths = np.array(val.select(monitor.tolist())['length']); labels = np.array(val.select(monitor.tolist())['labels'])
short = lengths <= 4096
for lr in p['lrs']:
run = {'name': f'pilot_lr{lr:g}', 'indices': 'pilot', 'rows': p['rows'], 'epochs': p['epochs'], 'lr': lr,
'alpha': p['alpha'], 'temperature': p['temperature'], 'start_from': 'init', 'evals_per_epoch': 1, 'export': False}
done = run_phase(run)
last = exports(run['name'])[-1]
logits = np.load(last/'monitor_logits.npz')['logits']
m_short = engine.metrics(labels[short], logits[short]); m_all = engine.metrics(labels, logits)
results.append({'lr': lr, 'short_nll': m_short['nll'], 'short_accuracy': m_short['accuracy'], 'short_balanced_accuracy': m_short['balanced_accuracy'],
'all_nll': m_all['nll'], 'all_accuracy': m_all['accuracy'], 'steps': done['steps'], 'export': str(last)})
event('pilot_result', **results[-1])
best = min(results, key=lambda r: r['short_nll'])
atomic_json(CKPT/'pilot_choice.json', {'lr': best['lr'], 'results': results, 'criterion': p['criterion']})
event('pilot_chosen', lr=best['lr'])
if __name__ == '__main__':
if sys.argv[1] == 'pilots': pilots()
else: run_phase(next(r for r in plan()['runs'] if r['name'] == sys.argv[1]))