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