"""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()