auto-200m-2 / training /engine.py
ProCreations's picture
auto-200m-2: ModernBERT-base approve/deny gate, 65,536-token context
0bbb929 verified
Raw History Blame Contribute Delete
17.6 kB
"""Full-parameter ModernBERT-base training with binary logit distillation (adapted from the auto-0.4b-2 iteration-1 engine)."""
import os
os.environ.setdefault('TOKENIZERS_PARALLELISM', 'false')
os.environ.setdefault('OMP_NUM_THREADS', '8')
import gc, hashlib, json, math, pathlib, random, shutil, signal, time
import numpy as np
import torch
import torch.nn.functional as F
from datasets import load_from_disk
from scipy.special import softmax
from sklearn.metrics import accuracy_score, f1_score, roc_auc_score, log_loss
from transformers import AutoModelForSequenceClassification, AutoTokenizer
from torch.utils.data import DataLoader
from common import ROOT, DATA, CKPT, LOG, ATTENTION, SEED, MAX_LEN, sha, event, atomic_json
# Gradient checkpointing only for microbatches above this many padded tokens (measured by smoke.py; conservative default).
CKPT_TOKENS = json.loads((DATA/'smoke.json').read_text())['ckpt_tokens'] if (DATA/'smoke.json').exists() else 16384
STOP = False
torch.set_num_threads(8)
torch.set_float32_matmul_precision('high')
torch.backends.cuda.matmul.allow_tf32 = True
def seed_all(seed=SEED):
random.seed(seed); np.random.seed(seed); torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)
def load_model(path, train=False):
model = AutoModelForSequenceClassification.from_pretrained(
str(path), dtype=torch.float32 if train else torch.bfloat16, attn_implementation=ATTENTION, allow_all_kernels=True).cuda()
assert model.config.id2label == {0: 'approve', 1: 'deny'}
assert model.config.max_position_embeddings == MAX_LEN
model.train(train)
return model
def metrics(labels, logits, threshold=0.5):
labels = np.asarray(labels)
probs = softmax(np.asarray(logits, dtype=np.float64), axis=1)[:, 1]
pred = probs >= threshold
deny, approve = labels == 1, labels == 0
recalls = ([float(pred[deny].mean())] if deny.any() else []) + ([float((~pred[approve]).mean())] if approve.any() else [])
out = {'n': len(labels), 'accuracy': float(accuracy_score(labels, pred)),
'balanced_accuracy': float(np.mean(recalls)),
'f1_deny': float(f1_score(labels, pred, zero_division=0)),
'auroc': float(roc_auc_score(labels, probs)) if deny.any() and approve.any() else None,
'nll': float(log_loss(labels, np.stack([1-probs, probs], axis=1), labels=[0, 1])),
'false_approve_rate': float((~pred[deny]).mean()) if deny.any() else None,
'false_deny_rate': float(pred[approve].mean()) if approve.any() else None,
'false_approve_count': int((~pred[deny]).sum()), 'false_deny_count': int(pred[approve].sum()),
'deny_count': int(deny.sum()), 'approve_count': int(approve.sum()), 'threshold': float(threshold),
'brier': float(np.mean((probs-labels)**2))}
p, n, z = out['accuracy'], len(labels), 1.959963984540054
center = (p + z*z/(2*n)) / (1+z*z/n)
half = z * math.sqrt(p*(1-p)/n + z*z/(4*n*n)) / (1+z*z/n)
out['accuracy_ci95'] = [center-half, center+half]
return out
def report(ds, indices, logits, threshold=0.5):
sub = ds.select([int(i) for i in indices])
labels = np.array(sub['labels'])
out = {'overall': metrics(labels, logits, threshold), 'slices': {}}
lens = np.array(sub['length'])
buckets = np.where(lens < 1024, '<1k', np.where(lens < 4096, '1k-4k', np.where(lens < 16384, '4k-16k', '16k-64k')))
for field, values in [('length', buckets), ('category', np.array(sub['category'])),
('difficulty', np.array(sub['difficulty'])), ('lang', np.array(sub['lang']))]:
out['slices'][field] = {str(v): metrics(labels[values == v], logits[values == v], threshold) for v in np.unique(values)}
return out
def collate(rows):
length = ((max(len(r['input_ids']) for r in rows)+7)//8)*8
ids = torch.full((len(rows), length), 50283, dtype=torch.long)
mask = torch.zeros_like(ids)
for i, row in enumerate(rows):
n = len(row['input_ids']); ids[i, :n] = torch.tensor(row['input_ids']); mask[i, :n] = 1
return {'input_ids': ids, 'attention_mask': mask, 'labels': torch.tensor([r['labels'] for r in rows], dtype=torch.long)}
def batches(indices, lengths, token_budget=16384, max_batch=32, seed=None):
indices = np.asarray(indices, dtype=np.int64).copy()
rng = np.random.default_rng(seed)
if seed is None:
indices = indices[np.argsort(lengths[indices], kind='stable')]
chunks = [indices]
else:
rng.shuffle(indices)
chunks = [c[np.argsort(lengths[c], kind='stable')] for c in np.array_split(indices, max(1, math.ceil(len(indices)/4096)))]
result = []
for chunk in chunks:
batch, maxlen = [], 0
for idx in chunk:
n = int(lengths[idx])
if batch and ((len(batch)+1) * max(maxlen, n) > token_budget or len(batch) >= max_batch):
result.append(batch); batch, maxlen = [], 0
batch.append(int(idx)); maxlen = max(maxlen, n)
if batch: result.append(batch)
if seed is not None: rng.shuffle(result)
return result
def optimizer_groups(microbatches, lengths, examples=128, tokens=131072):
groups, group, count, total = [], [], 0, 0
for batch in microbatches:
group.append(batch); count += len(batch); total += int(lengths[batch].sum())
if count >= examples or total >= tokens:
groups.append(group); group, count, total = [], 0, 0
if group: groups.append(group)
return groups
@torch.inference_mode()
def predict(model, ds, indices, name, output=None, token_budget=32768, max_batch=64):
# Same batching as the auto-0.4b-2 evaluations: BF16 logits of borderline items depend slightly on batch composition.
indices = np.asarray(indices, dtype=np.int64)
if output is not None and pathlib.Path(output).exists():
saved = np.load(output)
assert np.array_equal(saved['indices'], indices), 'Prediction cache index mismatch'
return saved['logits']
lengths = np.array(ds['length'])
bs = batches(indices, lengths, token_budget=token_budget, max_batch=max_batch)
loader = DataLoader(ds, batch_sampler=bs, collate_fn=collate, num_workers=4, pin_memory=True)
logits = np.full((len(ds), 2), np.nan, dtype=np.float32)
was_training = model.training; model.eval()
start, done, last = time.monotonic(), 0, 0
for batch_indices, batch in zip(bs, loader):
x = {k: v.cuda(non_blocking=True) for k, v in batch.items() if k != 'labels'}
with torch.autocast('cuda', dtype=torch.bfloat16):
pred = model(**x).logits.float().cpu().numpy()
if not np.isfinite(pred).all(): raise RuntimeError('Nonfinite evaluation logits')
logits[batch_indices] = pred; done += len(batch_indices)
if time.monotonic()-last > 45:
event('evaluating', name=name, done=done, total=len(indices), elapsed_seconds=time.monotonic()-start)
last = time.monotonic()
model.train(was_training)
result = logits[indices]
assert np.isfinite(result).all()
if output is not None:
pathlib.Path(output).parent.mkdir(parents=True, exist_ok=True)
temp = str(output) + '.tmp.npz'; np.savez(temp, indices=indices, logits=result); os.replace(temp, output)
event('evaluation_complete', name=name, n=len(indices), elapsed_seconds=time.monotonic()-start)
return result
def export_model(model, path, tokenizer):
path = pathlib.Path(path)
staging = path.with_name(path.name + '.staging')
if staging.exists(): shutil.rmtree(staging)
staging.mkdir(parents=True, exist_ok=True)
state = {k: v.detach().cpu().to(torch.bfloat16) if v.is_floating_point() else v.detach().cpu() for k,v in model.state_dict().items()}
old_dtype = model.config.dtype
model.config.dtype = torch.bfloat16
model.save_pretrained(str(staging), state_dict=state, safe_serialization=True)
model.config.dtype = old_dtype
tokenizer.save_pretrained(str(staging))
if path.exists(): shutil.rmtree(path)
staging.rename(path)
del state
def save_resume(model, optimizer, scheduler, path, **progress):
state = {'model': model.state_dict(), 'optimizer': optimizer.state_dict(), 'scheduler': scheduler.state_dict(),
'rng_torch': torch.get_rng_state(), 'rng_cuda': torch.cuda.get_rng_state_all(),
'rng_numpy': np.random.get_state(), 'rng_python': random.getstate(), **progress}
tmp = str(path) + '.tmp'; torch.save(state, tmp); os.replace(tmp, path)
def stop_handler(*_):
global STOP
STOP = True
def train_phase(name, start_path, indices, epochs, lr, teacher_logits=None, alpha=0.5, temperature=2.0,
token_budget=65536, max_batch=128, evals_per_epoch=4, seed_offset=0, min_lr_fraction=0.1, train_dir='train_all', deny_weight=1.0,
warmup_fraction=0.03, config=None, export=True):
global STOP
STOP = False
phase = CKPT / name; phase.mkdir(parents=True, exist_ok=True)
done_path = phase / 'complete.json'
if done_path.exists(): return json.loads(done_path.read_text())
seed_all(SEED + seed_offset)
train = load_from_disk(str(DATA / train_dir)); val = load_from_disk(str(DATA / 'validation'))
lengths = np.array(train['length'])
monitor = np.load(DATA / 'validation_partitions.npz')['monitor']
tokenizer = AutoTokenizer.from_pretrained(str(start_path))
model = load_model(start_path, train=True)
params = [p for p in model.parameters() if p.requires_grad]
optimizer = torch.optim.AdamW(params, lr=lr, betas=(0.9, 0.95), eps=1e-8, weight_decay=0.01, fused=True)
all_groups = [optimizer_groups(batches(indices, lengths, token_budget=token_budget, max_batch=max_batch, seed=SEED+seed_offset+e), lengths) for e in range(epochs)]
total_steps = sum(map(len, all_groups)); warmup = max(20, int(total_steps*warmup_fraction))
def schedule(step):
if step < warmup: return (step+1)/warmup
fraction = min(1., (step-warmup)/max(1,total_steps-warmup))
return min_lr_fraction + (1-min_lr_fraction)*0.5*(1+math.cos(math.pi*fraction))
scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, schedule)
start_epoch = start_group = global_step = 0
resume_path = phase / 'resume.pt'
signature = {'run_config': config, 'data_audit_sha256': sha(DATA/'data_audit.json'),
'teacher_logits_sha256': sha(teacher_logits) if teacher_logits else None,
'initial_weights_sha256': sha(pathlib.Path(start_path)/'model.safetensors'),
'indices_sha256': hashlib.sha256(np.asarray(indices, dtype=np.int64).tobytes()).hexdigest()}
history = []
if resume_path.exists():
resume = torch.load(resume_path, map_location='cpu', weights_only=False)
assert resume['signature'] == signature, 'Resume inputs or training plan changed'
model.load_state_dict(resume['model']); optimizer.load_state_dict(resume['optimizer']); scheduler.load_state_dict(resume['scheduler'])
start_epoch, start_group, global_step = resume['epoch'], resume['next_group'], resume['step']
history = resume.get('history', [])
torch.set_rng_state(resume['rng_torch']); torch.cuda.set_rng_state_all(resume['rng_cuda'])
np.random.set_state(resume['rng_numpy']); random.setstate(resume['rng_python'])
del resume
event('resumed', phase=name, epoch=start_epoch, group=start_group, step=global_step)
teacher = np.load(teacher_logits, mmap_mode='r') if teacher_logits else None
if teacher is not None: assert teacher.shape == (len(train), 2) and np.isfinite(teacher).all()
event('phase_started', phase=name, rows=len(indices), tokens=int(lengths[indices].sum()), epochs=epochs, lr=lr,
alpha=alpha, temperature=temperature, deny_weight=deny_weight, total_steps=total_steps, distillation=teacher is not None,
token_budget=token_budget, max_batch=max_batch)
last_log = last_save = time.monotonic(); started = last_log; seen = tokens_seen = 0; loss_sum = 0.; loss_examples = 0
ckpt_enabled = False
eval_interval = max(100, math.ceil(total_steps / (epochs*evals_per_epoch)))
signal.signal(signal.SIGTERM, stop_handler); signal.signal(signal.SIGINT, stop_handler)
for epoch, groups in enumerate(all_groups):
if epoch < start_epoch: continue
begin = start_group if epoch == start_epoch else 0
remaining_batches = [b for group in groups[begin:] for b in group]
loader = iter(DataLoader(train, batch_sampler=remaining_batches, collate_fn=collate, num_workers=6, pin_memory=True, prefetch_factor=4))
for group_idx in range(begin, len(groups)):
group = groups[group_idx]; n_group = sum(map(len, group)); optimizer.zero_grad(set_to_none=True)
for batch_indices in group:
batch = next(loader)
want_checkpoint = batch['input_ids'].numel() > CKPT_TOKENS
if want_checkpoint != ckpt_enabled:
if want_checkpoint: model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={'use_reentrant': False})
else: model.gradient_checkpointing_disable()
ckpt_enabled = want_checkpoint
x = {k: v.cuda(non_blocking=True) for k,v in batch.items()}
labels = x.pop('labels')
with torch.autocast('cuda', dtype=torch.bfloat16):
logits = model(**x).logits.float()
ce = F.cross_entropy(logits, labels, reduction='none')
if teacher is not None and alpha > 0:
tl = torch.tensor(np.array(teacher[batch_indices]), device='cuda', dtype=torch.float32)
target = F.softmax(tl/temperature, dim=-1)
kl = F.kl_div(F.log_softmax(logits/temperature, dim=-1), target, reduction='none').sum(-1)*temperature**2
# Uniform teacher weight: the teacher's judgement is trusted equally where it disagrees with the noisy label.
per_example = (1-alpha)*ce + alpha*kl
else: per_example = ce
if deny_weight != 1.0:
# Safety-weighted objective: errors on deny-labelled rows cost more than errors on approve-labelled rows.
per_example = per_example * torch.where(labels == 1, deny_weight, 1.0)
loss = per_example.sum()/n_group
if not torch.isfinite(loss): raise RuntimeError('Nonfinite training loss')
loss.backward()
loss_sum += float(per_example.detach().sum()); loss_examples += len(batch_indices)
seen += len(batch_indices); tokens_seen += int(lengths[batch_indices].sum())
del logits, loss, per_example, ce, x, labels
grad_norm = torch.nn.utils.clip_grad_norm_(params, 1.0, error_if_nonfinite=True)
optimizer.step(); scheduler.step(); global_step += 1
now = time.monotonic()
if now-last_log > 60 or global_step % 200 == 0:
event('training', phase=name, epoch=epoch+1, step=global_step, total_steps=total_steps,
loss=loss_sum/max(1,loss_examples), grad_norm=float(grad_norm), lr=scheduler.get_last_lr()[0],
examples_this_run=seen, tokens_per_second=tokens_seen/max(1,now-started), elapsed_seconds=now-started,
gpu_gb=torch.cuda.max_memory_allocated()/1e9)
loss_sum = 0.; loss_examples = 0; last_log=now
if global_step % eval_interval == 0 or group_idx == len(groups)-1:
pred = predict(model, val, monitor, name+f'-step{global_step}')
r = report(val, monitor, pred)
atomic_json(phase/f'monitor-{global_step}.json', r)
entry = {'step': global_step, 'accuracy': r['overall']['accuracy'], 'balanced_accuracy': r['overall']['balanced_accuracy'],
'nll': r['overall']['nll'], 'false_approve_count': r['overall']['false_approve_count'],
'false_deny_count': r['overall']['false_deny_count'], 'long_accuracy': r['slices']['length'].get('16k-64k', {}).get('accuracy')}
history.append(entry)
event('validation', phase=name, **entry)
if export or group_idx == len(groups)-1:
export_model(model, phase/f'step-{global_step}', tokenizer)
atomic_json(phase/f'step-{global_step}'/'training_progress.json',
{'run':name,'step':global_step,'signature':signature,'validation':r['overall'], 'lr': lr, 'alpha': alpha, 'temperature': temperature})
np.savez(phase/f'step-{global_step}'/'monitor_logits.npz', indices=monitor, logits=pred)
if now-last_save > 900 or STOP or group_idx == len(groups)-1:
save_resume(model, optimizer, scheduler, resume_path, epoch=epoch, next_group=group_idx+1, step=global_step, signature=signature, history=history)
last_save = time.monotonic()
if STOP:
event('stopped_safely', phase=name, step=global_step)
raise SystemExit(75)
atomic_json(done_path, {'steps':global_step,'history':history,'signature':signature})
event('phase_complete', phase=name, steps=global_step, history=history)
del optimizer, scheduler, params, model
gc.collect(); torch.cuda.empty_cache()
if resume_path.exists(): resume_path.unlink()
return json.loads(done_path.read_text())