"""8k -> 64k context for ModernBERT-base: choose the global-attention RoPE theta by masked-LM loss on training text (no labels, no validation/benchmark text), then build the sequence-classification initialization.""" import gc, sys, time import numpy as np import torch from datasets import load_from_disk from transformers import AutoConfig, AutoModelForMaskedLM, AutoModelForSequenceClassification, AutoTokenizer from common import * THETAS = [160000.0, 640000.0, 1280000.0, 2560000.0, 5120000.0, 10240000.0] BUCKETS = {'short': (256, 2048, 256), 'mid': (4096, 16384, 96), 'long': (16384, MAX_LEN, 64)} def config(theta): cfg = AutoConfig.from_pretrained(sources()['base']['path']) rp = getattr(cfg, 'rope_parameters', None) assert isinstance(rp, dict) and 'full_attention' in rp and 'sliding_attention' in rp, f'unexpected rope config {rp}' rp['full_attention']['rope_theta'] = float(theta) assert rp['sliding_attention']['rope_theta'] == 10000.0 cfg.rope_parameters = rp cfg.max_position_embeddings = MAX_LEN return cfg def check_theta(model, theta): """At least one rotary buffer must use the requested global theta (and one the unchanged local 10k).""" found = {} for name, buf in model.named_buffers(): if 'inv_freq' in name: dim = buf.numel()*2 for t in (theta, 10000.0): expected = 1.0/(t ** (torch.arange(0, dim, 2, dtype=torch.float32)/dim)) if torch.allclose(buf.float().cpu(), expected, rtol=1e-4): found[t] = name assert theta in found and 10000.0 in found, f'rotary buffers do not reflect theta {theta}: {found}' def samples(): ds = load_from_disk(str(DATA/'train')) lengths = np.array(ds['length']); rng = np.random.default_rng(SEED) out = {} for name, (lo, hi, n) in BUCKETS.items(): pool = np.flatnonzero((lengths >= lo) & (lengths < hi)) out[name] = [ds[int(i)]['input_ids'] for i in np.sort(rng.choice(pool, min(n, len(pool)), replace=False))] return out @torch.inference_mode() def mlm_loss(model, seqs, mask_id, seed): rng = np.random.default_rng(seed); total = 0.; count = 0 for ids in seqs: ids = np.array(ids); special = (ids == 50281) | (ids == 50282) m = (rng.random(len(ids)) < 0.15) & ~special x = ids.copy(); x[m] = mask_id inp = torch.tensor(x, device='cuda')[None] tgt = torch.tensor(ids, device='cuda'); mm = torch.tensor(m, device='cuda') with torch.autocast('cuda', dtype=torch.bfloat16): hidden = model.model(input_ids=inp, attention_mask=torch.ones_like(inp)).last_hidden_state[0] logits = model.decoder(model.head(hidden[mm])).float() # masked positions only (full logits at 50k tokens = 10 GB) total += float(torch.nn.functional.cross_entropy(logits, tgt[mm], reduction='sum')); count += int(m.sum()) return total/count def sweep(): dest = DATA/'rope_sweep.json' if dest.exists(): return json.loads(dest.read_text()) tok = AutoTokenizer.from_pretrained(sources()['base']['path']) seqs = samples(); results = [] for theta in THETAS: t0 = time.time() model = AutoModelForMaskedLM.from_pretrained(sources()['base']['path'], config=config(theta), dtype=torch.bfloat16, attn_implementation=ATTENTION, allow_all_kernels=True).cuda().eval() check_theta(model, theta) row = {'theta': theta, **{b: mlm_loss(model, s, tok.mask_token_id, SEED) for b, s in seqs.items()}, 'seconds': time.time()-t0} results.append(row); event('rope_sweep', **row) del model; gc.collect(); torch.cuda.empty_cache() base_short = results[0]['short'] eligible = [r for r in results if r['short'] <= 1.05*base_short] chosen = min(eligible, key=lambda r: r['mid'] + r['long']) out = {'results': results, 'chosen_theta': chosen['theta'], 'samples': {b: len(s) for b, s in seqs.items()}, 'rule': 'min(mid + long masked-LM loss) with short-context loss within 5% of the original 160k theta; 15% masking; training text only'} atomic_json(dest, out); event('rope_chosen', theta=chosen['theta']) return out def build_init(): dest = init_path() if (dest/'model.safetensors').exists(): print('init ready'); return theta = sweep()['chosen_theta'] cfg = config(theta) cfg.num_labels = 2; cfg.id2label = {0: 'approve', 1: 'deny'}; cfg.label2id = {'approve': 0, 'deny': 1} cfg.classifier_pooling = 'cls' # same pooling as the auto-0.4b lineage; CLS sits next to the proposed call cfg.architectures = ['ModernBertForSequenceClassification'] torch.manual_seed(SEED) model = AutoModelForSequenceClassification.from_pretrained(sources()['base']['path'], config=cfg, dtype=torch.float32) check_theta(model, theta) tok = AutoTokenizer.from_pretrained(sources()['base']['path']); tok.model_max_length = MAX_LEN stage = dest.with_name('init.staging') model.save_pretrained(str(stage), safe_serialization=True); tok.save_pretrained(str(stage)); stage.rename(dest) n = sum(p.numel() for p in model.parameters()) atomic_json(DATA/'init.json', {'theta': theta, 'parameters': n, 'base': sources()['base']}) event('init_ready', theta=theta, parameters=n) if __name__ == '__main__': build_init()