""" split_data.py ============= Performs stratified train/val/test split on the merged dataset. Split ratio: 70% train / 15% validation / 15% test Uses stratified splitting to maintain class distribution. Usage: python preprocessing/split_data.py """ import logging from pathlib import Path import pandas as pd import yaml from sklearn.model_selection import train_test_split # --------------------------------------------------------------------------- # Setup # --------------------------------------------------------------------------- logging.basicConfig( level=logging.INFO, format="%(asctime)s | %(levelname)s | %(message)s", ) logger = logging.getLogger(__name__) PROJECT_ROOT = Path(__file__).resolve().parent.parent PROCESSED_DIR = PROJECT_ROOT / "data" / "processed" # Load config with open(PROJECT_ROOT / "configs" / "config.yaml", "r") as f: CONFIG = yaml.safe_load(f) TEST_SIZE = CONFIG["data"]["test_size"] # 0.15 VAL_SIZE = CONFIG["data"]["val_size"] # 0.15 RANDOM_STATE = 42 def main(): # Load merged dataset merged_path = PROCESSED_DIR / "merged.csv" if not merged_path.exists(): logger.error(f"Merged dataset not found at {merged_path}") logger.error("Run merge_datasets.py first.") return df = pd.read_csv(merged_path) logger.info(f"Loaded merged dataset: {len(df):,} rows") # --- Ensure no data leakage: deduplicate by exact text --- before = len(df) df = df.drop_duplicates(subset=["text"], keep="first") logger.info(f"Dedup: {before:,} → {len(df):,}") # --- Stratified split --- # First split: train+val (85%) vs test (15%) train_val, test = train_test_split( df, test_size=TEST_SIZE, random_state=RANDOM_STATE, stratify=df["label"], ) # Second split: train (70%) vs val (15%) # val_size relative to train_val = 0.15 / 0.85 ≈ 0.1765 relative_val_size = VAL_SIZE / (1 - TEST_SIZE) train, val = train_test_split( train_val, test_size=relative_val_size, random_state=RANDOM_STATE, stratify=train_val["label"], ) # --- Save splits --- train_path = PROCESSED_DIR / "train.csv" val_path = PROCESSED_DIR / "val.csv" test_path = PROCESSED_DIR / "test.csv" train.to_csv(train_path, index=False) val.to_csv(val_path, index=False) test.to_csv(test_path, index=False) # --- Print statistics --- logger.info(f"\n{'=' * 70}") logger.info(f"Split results:") logger.info(f" Train: {len(train):>8,} rows ({len(train)/len(df)*100:.1f}%) → {train_path}") logger.info(f" Val: {len(val):>8,} rows ({len(val)/len(df)*100:.1f}%) → {val_path}") logger.info(f" Test: {len(test):>8,} rows ({len(test)/len(df)*100:.1f}%) → {test_path}") logger.info(f" Total: {len(df):>8,} rows") logger.info(f"\nClass distribution per split:") logger.info(f"{'Label':<25s} {'Train':>8s} {'Val':>8s} {'Test':>8s}") logger.info(f"{'-'*25} {'-'*8} {'-'*8} {'-'*8}") for label in sorted(df["label"].unique()): t_count = len(train[train["label"] == label]) v_count = len(val[val["label"] == label]) te_count = len(test[test["label"] == label]) logger.info(f"{label:<25s} {t_count:>8,} {v_count:>8,} {te_count:>8,}") logger.info(f"{'=' * 70}") # --- Verify no leakage --- train_texts = set(train["text"]) val_texts = set(val["text"]) test_texts = set(test["text"]) train_val_overlap = train_texts & val_texts train_test_overlap = train_texts & test_texts val_test_overlap = val_texts & test_texts if train_val_overlap or train_test_overlap or val_test_overlap: logger.error("DATA LEAKAGE DETECTED!") logger.error(f" Train-Val overlap: {len(train_val_overlap)}") logger.error(f" Train-Test overlap: {len(train_test_overlap)}") logger.error(f" Val-Test overlap: {len(val_test_overlap)}") else: logger.info("✅ No data leakage detected across splits.") if __name__ == "__main__": main()