Spaces:
Running on Zero
Running on Zero
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()
|