File size: 7,618 Bytes
29f25be | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 | """Prepare compact window indexes and manually check the actual data path."""
import argparse
import itertools
import json
import random
import time
import tomllib
from pathlib import Path
import numpy as np
from vimeml.training.data import (
IGNORE_INDEX, SPLITS, SentenceWindowDataset, collate_arrays,
make_loader, prepare_indexes, write_json,
)
ROOT = Path(__file__).resolve().parents[3]
def check_batch(batch, dataset):
"""Independently compare each batch row to its original complete sentence."""
inputs, labels, mask = batch["input_ids"], batch["labels"], batch["attention_mask"]
if inputs.dtype != np.int64 or labels.dtype != np.int64 or mask.dtype != np.bool_:
raise ValueError("Incorrect tensor dtypes.")
if inputs.ndim != 2 or inputs.shape != labels.shape or inputs.shape != mask.shape:
raise ValueError("Incorrect batch shapes.")
if not 1 <= inputs.shape[1] <= dataset.context_length:
raise ValueError("Batch exceeds model context.")
dataset._open()
pad = dataset._store.manifest["special_ids"]["pad"]
if np.any(labels[~mask] != IGNORE_INDEX) or np.any(inputs[~mask] != pad):
raise ValueError("Padding would contribute to loss.")
for row, length in enumerate(batch["lengths"]):
sentence = int(batch["sentence_index"][row])
start = int(batch["window_start"][row])
sequence = dataset._store[sentence]
expected_length = min(dataset.context_length, len(sequence) - 1 - start)
if int(length) != expected_length or start % dataset.context_length:
raise ValueError("Incorrect window boundary.")
if not np.array_equal(mask[row], np.arange(inputs.shape[1]) < length):
raise ValueError("Invalid attention mask.")
if (not np.array_equal(inputs[row, :length], sequence[start:start + length])
or not np.array_equal(labels[row, :length], sequence[start + 1:start + length + 1])):
raise ValueError("Incorrect next-token shift or sentence boundary crossing.")
if int(batch["source_id"][row]) != int(dataset._sources[sentence]):
raise ValueError("Incorrect source metadata.")
return int(mask.sum())
def check_core(dataset, seed):
rng = random.Random(seed)
indices = {0, len(dataset) - 1}
indices.update(rng.randrange(len(dataset)) for _ in range(32))
# Explicitly exercise every window of a deterministic selection of long sentences.
for position in range(min(12, len(dataset.sentences))):
indices.update(range(int(dataset.first[position]), int(dataset.end[position])))
try:
samples = [dataset[i] for i in sorted(indices)]
batch = collate_arrays(samples, pad_id=dataset._store.manifest["special_ids"]["pad"])
return {"checked_windows": len(samples), "checked_prediction_pairs": check_batch(batch, dataset)}
finally:
dataset.close()
def check_torch(dataset, config):
import torch
loader = make_loader(dataset, config["batch_size"], config["num_workers"], config["seed"])
pairs = batches = positions = rows = 0
begin = time.perf_counter()
iterator = iter(loader)
try:
for batch in itertools.islice(iterator, config["check_batches"]):
arrays = {key: value.numpy() for key, value in batch.items()}
pairs += check_batch(arrays, dataset)
positions += batch["input_ids"].numel()
rows += batch["input_ids"].shape[0]
batches += 1
if batches == 0:
raise ValueError("Empty DataLoader.")
finally:
# Releasing the owning DataLoader shuts down its persistent spawn workers.
del iterator, loader
dataset.close()
elapsed = time.perf_counter() - begin
return {"batches": batches, "windows": rows, "prediction_pairs": pairs,
"elapsed_seconds_including_worker_startup": elapsed,
"padding_fraction": 1 - pairs / positions,
"prediction_pairs_per_second_including_worker_startup": pairs / elapsed,
"workers": config["num_workers"]}
def check_gpu():
import torch
import torch.nn.functional as functional
result = {"torch_version": torch.__version__, "cuda_runtime": torch.version.cuda,
"cuda_available": torch.cuda.is_available()}
if result["cuda_available"]:
properties = torch.cuda.get_device_properties(0)
# Exercise forward/backward kernels; no model, checkpoint or training job.
values = torch.randn(16, 32, device="cuda", requires_grad=True)
weight = torch.randn(64, 32, device="cuda", requires_grad=True)
targets = torch.arange(16, device="cuda")
loss = functional.cross_entropy(functional.linear(values, weight), targets)
loss.backward()
torch.cuda.synchronize()
if not torch.isfinite(loss) or not torch.isfinite(weight.grad).all():
raise ValueError("CUDA forward/backward produced non-finite values.")
result.update(device=properties.name, total_memory_mib=properties.total_memory / 1024**2,
compute_capability=list(torch.cuda.get_device_capability(0)),
forward_backward_check="passed")
return result
def main(argv=None):
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--config", type=Path, default=ROOT / "configs/loader.toml")
parser.add_argument("--prepare-only", action="store_true", help="Index and NumPy checks only; PyTorch not required.")
parser.add_argument("--workers", type=int, help="Override DataLoader CPU process count.")
args = parser.parse_args(argv)
config = tomllib.loads(args.config.read_text(encoding="utf-8"))
if args.workers is not None:
config["num_workers"] = args.workers
if (config["context_length"] < 8 or config["context_length"] % 8 or
config["batch_size"] < 1 or config["num_workers"] < 0 or config["check_batches"] < 1):
parser.error("Use context_length divisible by 8, positive batch/check counts, nonnegative workers.")
if not args.prepare_only:
try:
import torch
except ImportError as error:
parser.error(f"Install requirements.txt; see docs/training.md for the CUDA 12.8 build: {error}")
token_dir = ROOT / config["token_dir"]
index_dir = ROOT / config["index_dir"]
manifest = prepare_indexes(token_dir, index_dir, config["context_length"])
report = {"mode": "prepare_only" if args.prepare_only else "pytorch_check",
"config": config, "index_splits": manifest["splits"], "core_checks": {}}
for split in SPLITS:
dataset = SentenceWindowDataset(token_dir, index_dir, split)
report["core_checks"][split] = check_core(dataset, config["seed"])
print(f"{split}: {len(dataset):,} windows; core check passed", flush=True)
if not args.prepare_only:
report["pytorch_checks"] = {}
# Test remains structurally checked above, never used for tuning/early stopping.
for split in ("train", "validation"):
dataset = SentenceWindowDataset(token_dir, index_dir, split)
report["pytorch_checks"][split] = check_torch(dataset, config)
print(f"{split}: PyTorch batches passed", flush=True)
report["environment"] = check_gpu()
report["status"] = "passed"
path = index_dir / ("prepare-check.json" if args.prepare_only else "loader-check.json")
write_json(path, report)
print(json.dumps(report, ensure_ascii=False, indent=2))
print(f"Report: {path}")
if __name__ == "__main__":
main()
|