import argparse from pathlib import Path import pandas as pd from sklearn.model_selection import train_test_split from utils import ensure_dir, save_json, set_seed def main() -> None: parser = argparse.ArgumentParser(description="Create stratified train/val/test splits.") parser.add_argument("--input", required=True) parser.add_argument("--output_dir", required=True) parser.add_argument("--train_ratio", type=float, default=0.7) parser.add_argument("--val_ratio", type=float, default=0.1) parser.add_argument("--test_ratio", type=float, default=0.2) parser.add_argument("--seed", type=int, default=42) args = parser.parse_args() total = args.train_ratio + args.val_ratio + args.test_ratio if abs(total - 1.0) > 1e-8: raise ValueError("train_ratio + val_ratio + test_ratio must equal 1.0") set_seed(args.seed) data = pd.read_csv(args.input) ensure_dir(args.output_dir) train_val, test = train_test_split( data, test_size=args.test_ratio, random_state=args.seed, stratify=data["label_id"], ) val_relative = args.val_ratio / (args.train_ratio + args.val_ratio) train, val = train_test_split( train_val, test_size=val_relative, random_state=args.seed, stratify=train_val["label_id"], ) splits = { "train": train.sort_values("node_id").reset_index(drop=True), "val": val.sort_values("node_id").reset_index(drop=True), "test": test.sort_values("node_id").reset_index(drop=True), } for name, frame in splits.items(): frame.to_csv(Path(args.output_dir) / f"{name}.csv", index=False) summary = { "seed": args.seed, "total_rows": int(len(data)), "ratios": { "train": args.train_ratio, "val": args.val_ratio, "test": args.test_ratio, }, "splits": { name: { "rows": int(len(frame)), "label_distribution": frame["label"].value_counts().to_dict(), } for name, frame in splits.items() }, } save_json(summary, Path(args.output_dir) / "split_summary.json") print("Saved train/val/test splits.") if __name__ == "__main__": main()