"""One focused check of real V2 crops, persistent workers and checkpoint resume.""" import contextlib import copy import io import json import signal import sys from pathlib import Path from unittest.mock import patch ROOT = Path(__file__).resolve().parents[2] sys.path.insert(0, str(ROOT / "src")) import numpy as np import torch from vimeml.training import train from vimeml.training.data import SentenceWindowDataset, make_loader, prepare_indexes, write_json from vimeml.training.data_v2 import PrefixCropWindowDataset def snapshots(loader): return [{key: value.clone() for key, value in batch.items()} for batch in loader] def identical(left, right): assert len(left) == len(right) for a, b in zip(left, right): assert a.keys() == b.keys() assert all(torch.equal(a[key], b[key]) for key in a) def fixture_resume(output): token_dir, index_dir = output / "fixture-tokens", output / "fixture-indexes" token_dir.mkdir() sequences = [ [2, *[4 + i % 60 for i in range(length)], 3] for length in (6, 8, 14, 15, 16, 17, 30, 41) ] offsets = np.cumsum([0, *map(len, sequences)], dtype=np.uint64) splits = {} for split in ("train", "validation", "test"): np.asarray([token for sequence in sequences for token in sequence], dtype=" 0 assert resumed["uncropped_prediction_pairs_seen"] == 2 * splits["train"]["prediction_pairs"] assert ( resumed["full_validation"]["last"]["prediction_pairs"] == splits["validation"]["prediction_pairs"] ) return { "exact_checkpoint_resume": True, "dropout": 0.1, "epochs": 2, "updates": resumed["step"], "interrupted_step": interrupted["step"], "trained_tokens": resumed["total_trained_tokens"], "dropped_tokens": resumed["prefix_crop_dropped_tokens"], "full_epoch_accounting": "passed", "validation_unaugmented": True, } def main(): torch.set_num_threads(2) output = ROOT / "outputs/model-checks/v2-prefix-crop" if output.exists() and any(output.iterdir()): raise FileExistsError("Use a new output for a new acceptance run.") output.mkdir(parents=True, exist_ok=True) token_dir, index_dir = ( ROOT / "artifacts/token-data/corpus-v2-16k", ROOT / "artifacts/training-data/corpus-v2-c128", ) original = SentenceWindowDataset(token_dir, index_dir) cropped = PrefixCropWindowDataset(token_dir, index_dir) forced = PrefixCropWindowDataset(token_dir, index_dir, probability=1.0) worker_loader = None try: indices = list(range(4096)) indices.extend(sorted({int(cropped.first[0]), int(cropped.first[0]) + 1, len(cropped) - 1})) counts = {} for epoch in (0, 1): cropped.set_epoch(epoch) lengths = cropped.window_lengths(indices) eligible = applied = dropped = 0 offsets = [] for index, length in zip(indices, lengths): before, after = original[index], cropped[(epoch, index)] offset = after["crop_offset"] assert len(after["input_ids"]) == int(length) assert after["input_ids"] == before["input_ids"][offset:] assert after["labels"] == before["labels"][offset:] assert after == cropped[(epoch, index)] if offset: assert before["window_start"] == 0 and len(after["labels"]) >= 8 assert after["input_ids"][0] != 2 if before["window_start"] > 0 or len(before["labels"]) <= 8: assert offset == 0 if before["window_start"] == 0 and len(before["labels"]) > 8: eligible += 1 force = forced[(epoch, index)] assert force["crop_offset"] > 0 and len(force["labels"]) >= 8 offsets.append(offset) applied += offset > 0 dropped += offset counts[str(epoch)] = { "samples": len(indices), "eligible": eligible, "cropped": applied, "crop_rate_among_eligible": applied / eligible, "dropped_tokens": dropped, } if epoch == 0: first_offsets = offsets else: assert first_offsets != offsets selected = list(range(256)) worker_loader = make_loader( cropped, batch_size=32, num_workers=2, indices=selected, shuffle=False, bucket_multiplier=4, ) for epoch in (0, 1): worker_loader.batch_sampler.set_epoch(epoch) actual = snapshots(worker_loader) reference_loader = make_loader( cropped, batch_size=32, num_workers=0, indices=selected, shuffle=False, bucket_multiplier=4, ) reference_loader.batch_sampler.set_epoch(epoch) reference = snapshots(reference_loader) identical(actual, reference) resume_loader = make_loader( cropped, batch_size=32, num_workers=0, indices=selected, shuffle=False, bucket_multiplier=4, start_batch=3, ) resume_loader.batch_sampler.set_epoch(epoch, start_batch=3) identical(actual[3:], snapshots(resume_loader)) del reference_loader, resume_loader for split in ("validation", "test"): unchanged = SentenceWindowDataset(token_dir, index_dir, split) try: sample = unchanged[0] assert "crop_offset" not in sample finally: unchanged.close() report = { "status": "passed", "probability": 0.30, "min_remaining_tokens": 8, "real_corpus_samples": counts, "scalar_vector_lengths_match": True, "same_epoch_reproducible": True, "different_epochs_change_crops": True, "persistent_worker_epochs": [0, 1], "worker_counts_compared": [0, 2], "prefetched_resume_batches_exact": True, "short_and_continuation_windows_unchanged": True, "validation_test_unchanged": True, "fixture_training": fixture_resume(output), "formal_training_started": False, } write_json(output / "report.json", report) print(json.dumps(report, ensure_ascii=False, indent=2)) finally: del worker_loader original.close() cropped.close() forced.close() if __name__ == "__main__": main()