Voltline's picture
Release VimeML V2.1 step40000 FP32 and Core ML INT8 (GPL-2.0)
29f25be verified
Raw History Blame Contribute Delete
7.62 kB
"""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()