xuan-luo/temp / utils /prepare_data.py
xuan-luo's picture
download
raw
8.21 kB
#!/usr/bin/env python3
"""Prepare GPT-2 tokenizer files and convert MLRA shards for Megatron-Core."""
from __future__ import annotations
import argparse
import json
import re
import shutil
import subprocess
import sys
from pathlib import Path
import numpy as np
ROOT = Path(__file__).resolve().parents[1]
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
from utils.paths import DELTAKV_ROOT
MLRA_MAGIC = 20240520
MLRA_VERSION = 1
MLRA_HEADER_BYTES = 256 * np.dtype(np.int32).itemsize
GPT2_VOCAB_SIZE = 50257
TRAIN_SHARD_RE = re.compile(r"fineweb_train_(\d{6})\.bin\Z")
VALID_SHARD_RE = re.compile(r"fineweb_val_(\d{6})\.bin\Z")
def _add_megatron_to_path(megatron_repo: Path) -> None:
if not (megatron_repo / "megatron").is_dir():
raise FileNotFoundError(
f"{megatron_repo} is not a Megatron-LM checkout; run env/bootstrap.sh first"
)
sys.path.insert(0, str(megatron_repo))
def _read_header(path: Path) -> int:
header = np.fromfile(path, dtype="<i4", count=256)
if header.size != 256:
raise ValueError(f"{path}: truncated MLRA header")
magic, version, token_count = map(int, header[:3])
if magic != MLRA_MAGIC or version != MLRA_VERSION:
raise ValueError(
f"{path}: expected magic/version {MLRA_MAGIC}/{MLRA_VERSION}, "
f"got {magic}/{version}"
)
if not 0 < token_count < 2**31:
raise ValueError(f"{path}: invalid token count {token_count}")
if np.any(header[3:] != 0):
raise ValueError(f"{path}: reserved MLRA header fields are non-zero")
expected_size = MLRA_HEADER_BYTES + token_count * np.dtype(np.uint16).itemsize
actual_size = path.stat().st_size
if actual_size != expected_size:
raise ValueError(f"{path}: expected {expected_size} bytes, got {actual_size}")
return token_count
def _validate_manifest(shards: list[Path], pattern: re.Pattern[str], first_index: int) -> None:
indices: list[int] = []
for shard in shards:
match = pattern.fullmatch(shard.name)
if match is None:
raise ValueError(f"Unexpected shard filename: {shard.name}")
indices.append(int(match.group(1)))
expected = list(range(first_index, first_index + len(indices)))
if indices != expected:
raise ValueError(
f"Shard indices must be unique and contiguous from {first_index:06d}; "
f"got {indices[:3]}...{indices[-3:]}"
)
def _convert_shards(
shards: list[Path],
output_prefix: Path,
overwrite: bool,
check_token_range: bool,
) -> dict[str, int]:
from megatron.core.datasets.indexed_dataset import IndexedDatasetBuilder
output_prefix.parent.mkdir(parents=True, exist_ok=True)
bin_path = output_prefix.with_suffix(".bin")
idx_path = output_prefix.with_suffix(".idx")
if (bin_path.exists() or idx_path.exists()) and not overwrite:
raise FileExistsError(
f"{output_prefix} already exists; pass --overwrite to recreate it"
)
bin_path.unlink(missing_ok=True)
idx_path.unlink(missing_ok=True)
builder = IndexedDatasetBuilder(str(bin_path), dtype=np.uint16)
total_tokens = 0
try:
for shard_index, shard in enumerate(shards, start=1):
token_count = _read_header(shard)
tokens = np.memmap(
shard,
mode="r",
dtype="<u2",
offset=MLRA_HEADER_BYTES,
shape=(token_count,),
)
if check_token_range:
maximum_token = int(tokens.max())
if maximum_token >= GPT2_VOCAB_SIZE:
raise ValueError(
f"{shard}: token {maximum_token} exceeds GPT-2 vocabulary "
f"maximum {GPT2_VOCAB_SIZE - 1}"
)
builder.add_document(tokens, [token_count])
total_tokens += token_count
print(
f"[{shard_index:04d}/{len(shards):04d}] {shard.name}: "
f"{token_count:,} tokens (total {total_tokens / 1e9:.3f}B)",
flush=True,
)
del tokens
builder.finalize(str(idx_path))
except BaseException:
try:
builder.data_file.close()
except Exception:
pass
bin_path.unlink(missing_ok=True)
idx_path.unlink(missing_ok=True)
raise
return {"shards": len(shards), "tokens": total_tokens}
def _download_tokenizer(output_dir: Path, overwrite: bool) -> None:
output_dir.mkdir(parents=True, exist_ok=True)
required = [output_dir / "vocab.json", output_dir / "merges.txt"]
if all(path.exists() for path in required) and not overwrite:
print(f"Tokenizer already present at {output_dir}")
return
hf = shutil.which("hf")
if hf is None:
raise RuntimeError(
"The `hf` CLI is required. Install it with "
"`curl -LsSf https://hf.co/cli/install.sh | bash -s`."
)
subprocess.run(
[
hf,
"download",
"openai-community/gpt2",
"--include",
"vocab.json",
"--include",
"merges.txt",
"--local-dir",
str(output_dir),
],
check=True,
)
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--source-dir",
type=Path,
default=DELTAKV_ROOT.parent / "DeltaIndexSharing" / "MLRA" / "data" / "fineweb-edu100B",
)
parser.add_argument(
"--output-dir",
type=Path,
default=DELTAKV_ROOT / "dataset" / "megatron",
)
parser.add_argument(
"--tokenizer-dir",
type=Path,
default=DELTAKV_ROOT / "dataset" / "tokenizer" / "gpt2",
)
parser.add_argument(
"--megatron-repo",
type=Path,
default=DELTAKV_ROOT / "env" / "third_party" / "Megatron-LM",
)
parser.add_argument("--max-train-shards", type=int)
parser.add_argument("--skip-tokenizer", action="store_true")
parser.add_argument("--skip-train", action="store_true")
parser.add_argument("--skip-valid", action="store_true")
parser.add_argument("--skip-token-range-check", action="store_true")
parser.add_argument("--overwrite", action="store_true")
args = parser.parse_args()
source_dir = args.source_dir.resolve()
if not source_dir.is_dir():
raise FileNotFoundError(f"MLRA data directory not found: {source_dir}")
if not args.skip_tokenizer:
_download_tokenizer(args.tokenizer_dir.resolve(), args.overwrite)
_add_megatron_to_path(args.megatron_repo.resolve())
metadata: dict[str, object] = {
"source_dir": str(source_dir),
"format": "Megatron IndexedDataset, uint16",
}
if not args.skip_train:
train_shards = sorted(source_dir.glob("fineweb_train_*.bin"))
_validate_manifest(train_shards, TRAIN_SHARD_RE, first_index=1)
if args.max_train_shards is not None:
train_shards = train_shards[: args.max_train_shards]
metadata["train"] = _convert_shards(
train_shards,
args.output_dir.resolve() / "fineweb_edu_100b_train",
args.overwrite,
not args.skip_token_range_check,
)
if not args.skip_valid:
valid_shards = sorted(source_dir.glob("fineweb_val_*.bin"))
if len(valid_shards) != 1:
raise ValueError(
f"Expected exactly one validation shard in {source_dir}, got {len(valid_shards)}"
)
_validate_manifest(valid_shards, VALID_SHARD_RE, first_index=0)
metadata["valid"] = _convert_shards(
valid_shards,
args.output_dir.resolve() / "fineweb_edu_100b_valid",
args.overwrite,
not args.skip_token_range_check,
)
metadata_path = args.output_dir.resolve() / "metadata.json"
metadata_path.parent.mkdir(parents=True, exist_ok=True)
metadata_path.write_text(json.dumps(metadata, indent=2) + "\n")
print(f"Wrote {metadata_path}")
if __name__ == "__main__":
main()

Xet Storage Details

Size:
8.21 kB
·
Xet hash:
8d6523134b39a448e19471e4bf3c1ae3a40c3f885c83a3de2e4a08ebb924a2dc

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.