auto-200m-2 / training /prepare.py
ProCreations's picture
auto-200m-2: ModernBERT-base approve/deny gate, 65,536-token context
0bbb929 verified
Raw History Blame Contribute Delete
10 kB
"""Full-length tokenization with the auto-0.4b-2 cleaning rules, verified against the cached Auto 3B logits;
plus the 120,000 archived cross-paired rows whose only label is the Auto 3B teacher."""
import os
os.environ['TOKENIZERS_PARALLELISM'] = 'false'
import collections
import hashlib
import json
import re
import shutil
import numpy as np
from datasets import Dataset, load_from_disk, concatenate_datasets, Value
from transformers import AutoTokenizer
from common import *
TOKENIZER = None
KEEP = ['label', 'category', 'difficulty', 'lang', 'text_hash', 'group_hash', 'source_row']
def digest(text):
return hashlib.sha256(re.sub(r'\s+', ' ', text).strip().encode()).hexdigest()
def fingerprints(batch, indices):
groups = []
for req, call in zip(batch['user_request'], batch['call']):
try: call = json.dumps(json.loads(call), sort_keys=True, ensure_ascii=False)
except (ValueError, TypeError): pass
groups.append(digest(req + '\n' + call))
return {'text_hash': [digest(t) for t in batch['text']], 'group_hash': groups, 'source_row': indices}
def raw(name, filename):
ds = Dataset.from_parquet(str(Path(sources()[name]['path'])/filename), cache_dir=str(DATA/'cache'))
assert set(ds.unique('label')) == {'approve', 'deny'}
return ds.map(fingerprints, with_indices=True, batched=True, batch_size=256, num_proc=8)
def clean(ds, forbidden_text, forbidden_group):
mapping, conflicts = {}, set()
for t, y in zip(ds['text_hash'], ds['label']):
if t in mapping and mapping[t] != y: conflicts.add(t)
mapping[t] = y
seen, keep, removed = set(), [], collections.Counter()
for i, (t, g) in enumerate(zip(ds['text_hash'], ds['group_hash'])):
reason = ('heldout_text' if t in forbidden_text else 'heldout_group' if g in forbidden_group else
'conflict' if t in conflicts else 'duplicate' if t in seen else None)
if reason: removed[reason] += 1
else: seen.add(t); keep.append(i)
return ds.select(keep), dict(removed)
def tokenize(batch):
global TOKENIZER
if TOKENIZER is None:
TOKENIZER = AutoTokenizer.from_pretrained(sources()['base']['path'])
ids = TOKENIZER(batch['text'], truncation=False, return_attention_mask=False)['input_ids']
lens = list(map(len, ids))
assert min(lens) > 0 and max(lens) <= MAX_LEN, 'Full context must fit; no silent truncation'
return {'input_ids': ids, 'length': lens, 'labels': [0 if y == 'approve' else 1 for y in batch['label']]}
def save(ds, name):
dest = DATA/name
if dest.exists(): return load_from_disk(str(dest))
out = ds.map(tokenize, batched=True, batch_size=48, num_proc=12,
remove_columns=[c for c in ds.column_names if c not in KEEP], desc=name)
stage = dest.with_name(dest.name+'.staging')
if stage.exists(): shutil.rmtree(stage)
out.save_to_disk(str(stage), max_shard_size='512MB'); stage.rename(dest)
return load_from_disk(str(dest))
def alignment(ds):
h = hashlib.sha256()
for t, row, label in zip(ds['text_hash'], ds['source_row'], ds['labels']):
h.update(f'{row}:{label}:{t}\n'.encode())
return h.hexdigest()
def check_tokenizer():
"""The base tokenizer must equal the one used by auto-0.4b-2 (ModernBERT-large) so token counts/lengths are comparable."""
a = AutoTokenizer.from_pretrained(sources()['base']['path']); b = AutoTokenizer.from_pretrained(sources()['reference']['path'])
assert a.get_vocab() == b.get_vocab()
for s in ['unicode Ω 中文 test', '### PROPOSED TOOL CALL\ntool: Bash\nargs: rm -rf /tmp/x', 'Привет, 世界 — ¿qué?']:
assert a(s)['input_ids'] == b(s)['input_ids']
assert (a.cls_token_id, a.sep_token_id, a.pad_token_id) == (50281, 50282, 50283)
def build_aug(val, bench, train):
"""Archived cross-paired rows: history+call of one training row judged under another row's request (same
language/framework/domain). Labels = Auto 3B argmax; soft targets = its logits. Rescreened against held-out data."""
raw_path = AUG/'aug_raw.parquet'; logits_path = AUG/'aug_teacher_logits.npy'
assert sha(logits_path) == AUG_LOGITS_SHA, 'augmentation teacher logits changed'
ds = Dataset.from_parquet(str(raw_path), cache_dir=str(DATA/'cache'))
tl = np.load(logits_path).astype(np.float32)
assert tl.shape == (len(ds), 2) and np.isfinite(tl).all() and len(ds) == 120000
ds = ds.map(lambda b: {'group_hash': fingerprints(b, [0]*len(b['text']))['group_hash'],
'text_hash2': [digest(t) for t in b['text']]}, batched=True, batch_size=256, num_proc=8)
assert list(ds['text_hash']) == list(ds['text_hash2'])
forbidden_text = set(val['text_hash']) | set(bench['text_hash']) | set(train['text_hash'])
forbidden_group = set(val['group_hash']) | set(bench['group_hash'])
keep, removed, seen = [], collections.Counter(), set()
for i, (t, g) in enumerate(zip(ds['text_hash'], ds['group_hash'])):
reason = 'heldout_or_train_text' if t in forbidden_text else 'heldout_group' if g in forbidden_group else 'duplicate' if t in seen else None
if reason: removed[reason] += 1
else: keep.append(i); seen.add(t)
labels = tl.argmax(1)
ds = ds.select(keep); tl = tl[keep]; labels = labels[keep]
ds = ds.map(lambda b, idx: {'label': ['deny' if labels[i] else 'approve' for i in idx], 'source_row': [1_000_000 + keep[i] for i in idx]},
with_indices=True, batched=True, batch_size=1024)
ds = ds.remove_columns([c for c in ds.column_names if c not in KEEP + ['text']])
return ds, tl, dict(removed)
def main():
if (DATA/'data_audit.json').exists(): print('data ready'); return
event('preparing_data')
check_tokenizer()
bench = raw('benchmark', 'test.parquet')
bt, bg = set(bench['text_hash']), set(bench['group_hash'])
val, vr = clean(raw('data', 'validation.parquet'), bt, bg)
vt, vg = set(val['text_hash']), set(val['group_hash'])
train, tr = clean(raw('data', 'train.parquet'), bt|vt, bg|vg)
assert len(train) == 711985 and len(val) == 13000 and len(bench) == 3000
student = save(train, 'train')
align = alignment(student)
receipt = json.loads((TEACHER_LOGITS/'teacher_logits.complete.json').read_text())
assert receipt['metadata']['alignment_sha256'] == align, 'Training rows differ from the cached teacher targets'
assert receipt['metadata']['teacher_sha256'] == TEACHER_SHA and receipt['metadata']['rows'] == len(student)
if not (DATA/'teacher_logits_train.npy').exists():
shutil.copy2(TEACHER_LOGITS/'teacher_logits.npy', DATA/'teacher_logits_train.npy.tmp')
os.replace(DATA/'teacher_logits_train.npy.tmp', DATA/'teacher_logits_train.npy')
assert sha(DATA/'teacher_logits_train.npy') == receipt['sha256']
tl = np.load(DATA/'teacher_logits_train.npy'); assert tl.shape == (len(student), 2) and np.isfinite(tl).all()
validation = save(val, 'validation')
benchmark = save(bench, 'benchmark')
partition = np.array([int(h[:8], 16) % 10 for h in validation['group_hash']])
parts = {'selection': np.flatnonzero(partition < 6), 'calibration': np.flatnonzero((partition >= 6) & (partition < 8)),
'audit': np.flatnonzero(partition >= 8)}
assert [len(parts[k]) for k in ['selection', 'calibration', 'audit']] == [7824, 2581, 2595]
archived = np.load(AUG/'validation_partitions.npz')
for k in ['selection', 'calibration', 'audit']: assert np.array_equal(archived[k], parts[k]), 'validation partitions differ from auto-0.4b-2'
parts['monitor'] = parts['selection']
np.savez(DATA/'validation_partitions.npz', **parts)
# Cross-paired teacher-labelled rows, tokenized with the student tokenizer, appended after the original rows.
aug_raw, aug_tl, ar = build_aug(val, bench, train)
aug = save(aug_raw, 'aug')
cols = ['input_ids', 'length', 'labels', 'label', 'category', 'difficulty', 'lang', 'text_hash', 'group_hash', 'source_row']
orig = student.select_columns(cols); aug = aug.select_columns(cols).cast(orig.features)
combined = concatenate_datasets([orig, aug])
if not (DATA/'train_all').exists():
if (DATA/'train_all.staging').exists(): shutil.rmtree(DATA/'train_all.staging')
combined.save_to_disk(str(DATA/'train_all.staging'), max_shard_size='512MB'); (DATA/'train_all.staging').rename(DATA/'train_all')
tl_all = np.concatenate([tl, aug_tl], axis=0).astype(np.float32)
assert len(tl_all) == len(combined)
np.save(DATA/'teacher_logits_all.npy', tl_all)
labels = np.array(student['labels'])
summary = {'removed_train': tr, 'removed_validation': vr, 'removed_augmented': ar, 'alignment_sha256': align,
'no_truncation': True, 'teacher_target_split': 'train only', 'teacher_weights_sha256': TEACHER_SHA,
'teacher_logits_sha256': receipt['sha256'], 'teacher_train_agreement': float((tl.argmax(1) == labels).mean()),
'augmented_rows': len(aug), 'augmented_teacher_deny_fraction': float(aug_tl.argmax(1).mean()),
'teacher_logits_all_sha256': sha(DATA/'teacher_logits_all.npy'),
'validation_partitions': {k: len(v) for k, v in parts.items()}}
for name, ds in [('train', student), ('augmented', aug), ('validation', validation), ('benchmark', benchmark)]:
lens = np.array(ds['length'])
summary[name] = {'rows': len(ds), 'tokens': int(lens.sum()), 'max_length': int(lens.max()),
'gt4096': int((lens > 4096).sum()), 'gt16384': int((lens > 16384).sum()),
'tokens_le4096': int(lens[lens <= 4096].sum()), 'tokens_gt4096': int(lens[lens > 4096].sum()),
'labels': dict(collections.Counter(ds['label']))}
assert summary['train']['tokens'] == 516733212 and summary['benchmark']['tokens'] == 11710540, 'token counts differ from auto-0.4b-2'
atomic_json(DATA/'data_audit.json', summary)
event('data_ready', summary=summary)
if __name__ == '__main__': main()