"""Step 2 — padding ratio. Data-side, no GPU, no model. Replicates the exact batching the trainer used: SequentialSampler over the encode-cache order, per_device_batch_size=2, collate_mixed right-padding to the pair max. Reports real vs padded token counts and the loss-target count that sizes the chunked CE. """ import json, sys from collections import Counter CACHE = "/tank/erp-tune/run-01/encode-cache/encoded-a4b0796de1260930.jsonl" IGNORE_INDEX = -100 MB = 2 # per_device_batch_size ACCUM = 8 # gradient_accumulation_steps lens, kept_counts, kinds = [], [], [] with open(CACHE) as fh: for line in fh: row = json.loads(line) ids = row["input_ids"] labels = row["labels"] lens.append(len(ids)) kept_counts.append(sum(1 for x in labels if x != IGNORE_INDEX)) kinds.append(row.get("sample_kind", "?")) n = len(lens) print(f"records {n:,}") print(f"sample_kind mix {dict(Counter(kinds))}") print() print(f"seq len min/mean/max {min(lens)} / {sum(lens)/n:.0f} / {max(lens)}") print(f"loss targets min/mean/max {min(kept_counts)} / {sum(kept_counts)/n:.0f} / {max(kept_counts)}") print() # --- micro-batch padding, exactly as collate_mixed builds it --- real = padded = 0 mb_widths, mb_waste, mb_kept = [], [], [] for i in range(0, n - n % MB, MB): group = lens[i:i + MB] width = max(group) r = sum(group) p = width * MB real += r padded += p mb_widths.append(width) mb_waste.append(1 - r / p) mb_kept.append(sum(kept_counts[i:i + MB])) nb = len(mb_widths) print(f"micro-batches (mb={MB}) {nb:,}") print(f"real tokens {real:,}") print(f"padded tokens {padded:,}") print(f"PADDING WASTE {100 * (1 - real / padded):.1f}% ({padded - real:,} pad tokens)") print() print(f"mb width min/mean/max {min(mb_widths)} / {sum(mb_widths)/nb:.0f} / {max(mb_widths)}") srt = sorted(mb_widths) for q in (50, 75, 90, 95, 99): print(f" p{q} width {srt[int(nb*q/100)]}") print(f"mb at max_seq_len 16384 {sum(1 for w in mb_widths if w >= 16384):,} ({100*sum(1 for w in mb_widths if w>=16384)/nb:.1f}%)") print() srtw = sorted(mb_waste) print(f"per-mb waste p50/p90/max {100*srtw[nb//2]:.1f}% / {100*srtw[int(nb*0.9)]:.1f}% / {100*max(mb_waste):.1f}%") print() print(f"loss targets per mb min/mean/max {min(mb_kept)} / {sum(mb_kept)/nb:.0f} / {max(mb_kept)}") print(f" -> CE chunks per mb (1024) min/mean/max {min(mb_kept)//1024+1} / {sum(mb_kept)/nb/1024:.1f} / {max(mb_kept)//1024+1}") print() # --- what length-bucketing would recover (sort by length, then batch) --- order = sorted(range(n), key=lambda i: lens[i]) b_real = b_padded = 0 for i in range(0, n - n % MB, MB): group = [lens[j] for j in order[i:i + MB]] b_real += sum(group) b_padded += max(group) * MB print("--- counterfactual: length-bucketed sampler ---") print(f"bucketed padded tokens {b_padded:,}") print(f"bucketed waste {100 * (1 - b_real / b_padded):.1f}%") print(f"TOKEN REDUCTION vs current {100 * (1 - b_padded / padded):.1f}%")