"""Measure bucket-to-pair / shuffle-to-mix against the REAL encode cache. Brokkr's design, validated on measured record lengths rather than a calibrated length model: 1. sort records by length 2. cut into buckets of BUCKET records 3. form micro-batches of 2 WITHIN each bucket (adjacent after sort) 4. shuffle the resulting MICRO-BATCHES globally, seeded Padding efficiency is a property of the pairing only, so step 4 costs nothing and restores root-mixing inside each accumulation window. Also applies the fitted cost model from the replica scaling test to convert token savings into predicted wall clock. """ import json import random from collections import Counter CACHE = "/tank/erp-tune/run-01/encode-cache/encoded-a4b0796de1260930.jsonl" IGNORE_INDEX = -100 MB = 2 ACCUM = 8 SEED = 20260824 # fitted on the replica: t(w) = A*w + B*w^2 for a batch of 2 sequences of len w A = 6.8715e-04 B = 8.8509e-08 rows = [] with open(CACHE) as fh: for line in fh: r = json.loads(line) rows.append((len(r["input_ids"]), r.get("dataset_id", "?"), r.get("sample_kind", "?"))) n = len(rows) print("records %d" % n) print() def evaluate(order, label, show_roots=False): real = padded = 0 widths = [] batches = [] for i in range(0, n - n % MB, MB): grp = [rows[j] for j in order[i:i + MB]] w = max(g[0] for g in grp) real += sum(g[0] for g in grp) padded += w * MB widths.append(w) batches.append([g[1] for g in grp]) nb = len(widths) Ew = sum(widths) / nb Ew2 = sum(w * w for w in widths) / nb t_mb = A * Ew + B * Ew2 srt = sorted(widths) print("--- %s ---" % label) print(" padded tokens %s" % f"{padded:,}") print(" waste %.1f%%" % (100 * (1 - real / padded))) print(" E[w] (per-seq) %.0f" % Ew) print(" E[w^2] %.3e" % Ew2) print(" width p50/p90/p99 %d / %d / %d" % ( srt[nb // 2], srt[int(nb * .9)], srt[int(nb * .99)])) print(" predicted micro-batch %.3f s (lin %.3f + quad %.3f, quad %.0f%%)" % ( t_mb, A * Ew, B * Ew2, 100 * B * Ew2 / t_mb)) print(" predicted step (x%d) %.1f s -> %.2f h over 1312 steps" % ( ACCUM, t_mb * ACCUM, t_mb * ACCUM * 1312 / 3600)) # unpadded micro-batches take the is_causal fast path on the 5 global layers exact = sum(1 for i in range(0, n - n % MB, MB) if len(set(rows[j][0] for j in order[i:i + MB])) == 1) print(" ZERO-PAD micro-batches %d / %d (%.1f%%) <- global layers on is_causal" % ( exact, nb, 100 * exact / nb)) if show_roots: # root diversity inside an accumulation window div = [] for i in range(0, nb - nb % ACCUM, ACCUM): win = [d for b in batches[i:i + ACCUM] for d in b] div.append(len(set(win))) print(" roots per accum window mean %.2f min %d (of %d roots)" % ( sum(div) / len(div), min(div), len({r[1] for r in rows}))) print() return padded, t_mb # --- current: encode-cache order, SequentialSampler --- cur_padded, cur_t = evaluate(list(range(n)), "CURRENT (SequentialSampler)", True) # --- bucket-to-pair + shuffle-to-mix --- # BUCKET controls the efficiency-vs-diversity trade: records are globally # sorted, cut into buckets of BUCKET, SHUFFLED WITHIN the bucket (not # re-sorted), then paired adjacently. BUCKET=2 is a perfect global sort # (0% waste, worst root mixing); larger buckets admit more length spread # inside a pair but draw partners from a wider slice of the corpus. for BUCKET in (2, 8, 32, 128, 512): by_len = sorted(range(n), key=lambda i: rows[i][0]) rng = random.Random(SEED) micro = [] for s in range(0, n, BUCKET): chunk = by_len[s:s + BUCKET] rng.shuffle(chunk) # mix WITHIN the length bucket for k in range(0, len(chunk) - len(chunk) % MB, MB): micro.append(chunk[k:k + MB]) rng.shuffle(micro) # shuffle-to-mix across buckets order = [i for b in micro for i in b] placed = set(order) order += [i for i in by_len if i not in placed] p, t = evaluate(order, "BUCKET=%d, shuffle within + global micro-batch shuffle" % BUCKET, True) print(" >>> vs current: %.1f%% fewer padded tokens, %.1f%% less wall clock" % ( 100 * (1 - p / cur_padded), 100 * (1 - t / cur_t))) print()