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