File size: 2,288 Bytes
877049d | 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 | 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()
|