Spaces:
Running on Zero
Running on Zero
Download preprocessing/split_data.py from heatherlead/comment_moderation: direct link, hf CLI and curl.
- Browser
- Download file 4.08 kB
-
https://huggingface.co/spaces/heatherlead/comment_moderation/resolve/main/preprocessing/split_data.py
- Command line
-
hf download hf://spaces/heatherlead/comment_moderation/preprocessing/split_data.py
-
curl -L -o split_data.py https://huggingface.co/spaces/heatherlead/comment_moderation/resolve/main/preprocessing/split_data.py
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() | |