heatherlead's picture
Add moderation backend with LFS tracking for large files
26c8f44
Raw History Blame Contribute Delete
4.08 kB
"""
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()