infosec-v1 / code /training /scripts /split_jsonl.py
adhikjoshi's picture
Super-squash branch 'main' using huggingface_hub
994182c
Raw
History Blame Contribute Delete
1.74 kB
#!/usr/bin/env python3
"""Deterministically split normalized JSONL into train/validation files."""
from __future__ import annotations
import argparse
import hashlib
from pathlib import Path
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--input", required=True)
parser.add_argument("--train", required=True)
parser.add_argument("--val", required=True)
parser.add_argument("--val-ratio", type=float, default=0.02)
parser.add_argument("--seed", type=int, default=1337)
return parser.parse_args()
def score(line: str, seed: int) -> float:
digest = hashlib.sha256(f"{seed}:{line}".encode("utf-8")).digest()
value = int.from_bytes(digest[:8], "big")
return value / float(2**64 - 1)
def main() -> int:
args = parse_args()
if not 0 < args.val_ratio < 1:
raise ValueError("--val-ratio must be between 0 and 1")
input_path = Path(args.input)
train_path = Path(args.train)
val_path = Path(args.val)
train_path.parent.mkdir(parents=True, exist_ok=True)
val_path.parent.mkdir(parents=True, exist_ok=True)
train_count = 0
val_count = 0
with input_path.open("r", encoding="utf-8") as src, train_path.open("w", encoding="utf-8") as train, val_path.open(
"w", encoding="utf-8"
) as val:
for line in src:
if not line.strip():
continue
if score(line, args.seed) < args.val_ratio:
val.write(line)
val_count += 1
else:
train.write(line)
train_count += 1
print(f"train={train_count} val={val_count}")
return 0
if __name__ == "__main__":
raise SystemExit(main())