import argparse
import sys
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parents[1]))

import pandas as pd
from sklearn.model_selection import train_test_split


def main() -> int:
    parser = argparse.ArgumentParser(description="Create train/valid/test split definition CSV.")
    parser.add_argument("--manifest", required=True)
    parser.add_argument("--output", default="data/manifests/02_split_definition.csv")
    parser.add_argument("--valid-size", type=float, default=0.15)
    parser.add_argument("--test-size", type=float, default=0.15)
    parser.add_argument("--seed", type=int, default=42)
    args = parser.parse_args()

    df = pd.read_csv(args.manifest)
    train_df, temp_df = train_test_split(df, test_size=args.valid_size + args.test_size, random_state=args.seed, stratify=df["label_id"])
    valid_ratio = args.valid_size / (args.valid_size + args.test_size)
    valid_df, test_df = train_test_split(temp_df, test_size=1 - valid_ratio, random_state=args.seed, stratify=temp_df["label_id"])
    out = pd.concat(
        [
            train_df.assign(split="train"),
            valid_df.assign(split="valid"),
            test_df.assign(split="test"),
        ]
    )[["image_id", "split"]]
    out["reason"] = "stratified_by_label_id"
    out.to_csv(args.output, index=False)
    print(f"created: {args.output}")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
