File size: 4,079 Bytes
26c8f44
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
"""
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()