"""Reproducible full-parameter ModernBERT fine-tuning and binary logit distillation.""" 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, balanced_accuracy_score, f1_score, roc_auc_score, log_loss from transformers import AutoModelForSequenceClassification, AutoTokenizer from torch.utils.data import DataLoader ROOT = pathlib.Path('/home/user/.local/share/rtx-pro-apps/auto-0.4b-2') DATA = pathlib.Path('/home/user/datasets/auto-0.4b-2') CKPT = pathlib.Path('/home/user/checkpoints/auto-0.4b-2') LOG = pathlib.Path('/home/user/logs/auto-0.4b-2') ATTENTION = 'kernels-community/flash-attn2@81fb77c12b2ad5d69380669b46739d5868614502' SEED = 20260908 STOP = False torch.set_num_threads(8) torch.set_float32_matmul_precision('high') torch.backends.cuda.matmul.allow_tf32 = True def atomic_json(path, value): path = pathlib.Path(path) path.parent.mkdir(parents=True, exist_ok=True) tmp = path.with_suffix(path.suffix + '.tmp') tmp.write_text(json.dumps(value, indent=2, allow_nan=False)) tmp.replace(path) def event(kind, **kwargs): item = {'event': kind, 'time': time.strftime('%Y-%m-%dT%H:%M:%SZ', time.gmtime()), **kwargs} print(json.dumps(item, allow_nan=False), flush=True) LOG.mkdir(exist_ok=True, parents=True) with (LOG / 'events.jsonl').open('a') as f: f.write(json.dumps(item, allow_nan=False) + '\n') atomic_json(ROOT / 'status.json', item) 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).cuda() assert model.config.id2label == {0: 'approve', 1: 'deny'} assert model.config.max_position_embeddings == 65536 model.train(train) return model def release(model): del model gc.collect(); torch.cuda.empty_cache() 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))} # Wilson binomial interval for accuracy. 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)/2048)))] 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): 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=64) loader = DataLoader(ds, batch_sampler=bs, collate_fn=collate, num_workers=2, 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 selection_score(r): short_and_long = r['overall']['balanced_accuracy'] long = r['slices']['length'].get('16k-64k', {}).get('balanced_accuracy', short_and_long) return 0.7 * short_and_long + 0.3 * long def export_model(model, path, tokenizer): path = pathlib.Path(path) staging = path.with_name(path.name + '.staging') backup = path.with_name(path.name + '.previous') 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 backup.exists(): shutil.rmtree(backup) if path.exists(): path.rename(backup) staging.rename(path) if backup.exists(): shutil.rmtree(backup) 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): 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())['best_path'] seed_all() train = load_from_disk(str(DATA / 'train')); 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, seed=SEED+e), lengths) for e in range(epochs)] total_steps = sum(map(len, all_groups)); warmup = max(20, int(total_steps*0.03)) def schedule(step): if step < warmup: return (step+1)/warmup fraction = min(1., (step-warmup)/max(1,total_steps-warmup)) return 0.1 + 0.9*0.5*(1+math.cos(math.pi*fraction)) scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, schedule) start_epoch = start_group = global_step = 0; best_score = -1.; best_path = str(phase/'best') resume_path = phase / 'resume.pt' if resume_path.exists(): resume = torch.load(resume_path, map_location='cpu', weights_only=False) 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'] best_score = resume['best_score'] if (phase/'best_metrics.json').exists(): best_score = max(best_score, json.loads((phase/'best_metrics.json').read_text())['score']) 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) else: initial = predict(model, val, monitor, name+'-initial') initial_report = report(val, monitor, initial); best_score = selection_score(initial_report) export_model(model, best_path, tokenizer) atomic_json(phase/'best_metrics.json', {'score': best_score, 'step': 0, 'report': initial_report}) 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), epochs=epochs, lr=lr, total_steps=total_steps, distillation=teacher is not None) 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*4))) 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 # Background CPU collation, in exact deterministic group order. 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=4, pin_memory=True, prefetch_factor=2)) 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'].shape[1] > 8192 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: 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 # Preserve hard-label evidence on teacher mistakes while still learning soft uncertainty. weight = torch.where(tl.argmax(-1) == labels, alpha, alpha*0.25) per_example = (1-weight)*ce + weight*kl else: per_example = ce 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 > 45 or global_step % 100 == 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) 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); score=selection_score(r) atomic_json(phase/f'monitor-{global_step}.json', r) event('validation', phase=name, step=global_step, score=score, metrics=r['overall'], long_metrics=r['slices']['length'].get('16k-64k')) if score > best_score: best_score = score; export_model(model, best_path, tokenizer) atomic_json(phase/'best_metrics.json', {'score':score,'step':global_step,'report':r}) if now-last_save > 600 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, best_score=best_score) last_save = time.monotonic() if STOP: event('stopped_safely', phase=name, step=global_step) raise SystemExit(75) export_model(model, phase/'last', tokenizer) atomic_json(done_path, {'best_path':best_path,'score':best_score,'steps':global_step,'last_path':str(phase/'last')}) event('phase_complete', phase=name, steps=global_step, best_score=best_score) del optimizer, scheduler, params, model gc.collect(); torch.cuda.empty_cache() return best_path def baseline(): sources = json.loads((ROOT/'sources.json').read_text()) ds = load_from_disk(str(DATA/'benchmark')); val = load_from_disk(str(DATA/'validation')) selection = np.load(DATA/'validation_partitions.npz')['selection'] out = CKPT/'baseline'; out.mkdir(exist_ok=True, parents=True) student_tok=AutoTokenizer.from_pretrained(sources['student']['path']); teacher_tok=AutoTokenizer.from_pretrained(sources['teacher']['path']) assert student_tok.get_vocab() == teacher_tok.get_vocab() assert student_tok('unicode Ω 中文 test')['input_ids'] == teacher_tok('unicode Ω 中文 test')['input_ids'] for name in ['student','teacher']: if (out/f'{name}.json').exists(): continue model=load_model(sources[name]['path']) logits=predict(model,ds,np.arange(len(ds)),name+'-benchmark',out/f'{name}-benchmark.npz') r={'benchmark':report(ds,np.arange(len(ds)),logits)} logits=predict(model,val,selection,name+'-validation',out/f'{name}-selection.npz') r['selection']=report(val,selection,logits) atomic_json(out/f'{name}.json',r) event('baseline_complete',model=name,benchmark=r['benchmark']['overall'],selection=r['selection']['overall']) del model; gc.collect(); torch.cuda.empty_cache() if __name__ == '__main__': import sys if sys.argv[1] == 'baseline': baseline()