Skip to content

Splitting

Deterministic 10-fold splits with patient-level grouping. Folds are 1-indexed, and the default split mapping is derived from the fold count rather than fixed: train=1..n-2, val=n-1, test=n.

Engine

engine

Universal splitting engine.

Dispatches to StratifiedGroupKFold (patient-aware), StratifiedKFold, or reads predefined splits based on the dataset config.

split_dataset

split_dataset(df: DataFrame, labels: Series, config: DatasetConfig, n_folds: int = 10, random_state: int = 42) -> SplitResult

Universal splitting pipeline.

Decision logic: 1. If config.has_predefined_splits: read fold assignments from config 2. Else if config.patient_id_column is set: StratifiedGroupKFold 3. Else: StratifiedKFold

Parameters:

Name Type Description Default
df DataFrame

Metadata DataFrame (from splitter.load_metadata)

required
labels Series

Stratification labels (from splitter.get_stratification_labels)

required
config DatasetConfig

The dataset's DatasetConfig

required
n_folds int

Number of folds (default 10)

10
random_state int

Random seed for determinism

42

Returns:

Type Description
SplitResult

SplitResult with fold assignments (1-indexed folds)

Source code in ecgbench/splitting/engine.py
def split_dataset(
    df: pd.DataFrame,
    labels: pd.Series,
    config: DatasetConfig,
    n_folds: int = 10,
    random_state: int = 42,
) -> SplitResult:
    """Universal splitting pipeline.

    Decision logic:
    1. If config.has_predefined_splits: read fold assignments from config
    2. Else if config.patient_id_column is set: StratifiedGroupKFold
    3. Else: StratifiedKFold

    Args:
        df: Metadata DataFrame (from splitter.load_metadata)
        labels: Stratification labels (from splitter.get_stratification_labels)
        config: The dataset's DatasetConfig
        n_folds: Number of folds (default 10)
        random_state: Random seed for determinism

    Returns:
        SplitResult with fold assignments (1-indexed folds)
    """
    if config.has_predefined_splits and config.predefined_splits:
        return _split_predefined(df, labels, config)
    elif config.patient_id_column and config.patient_id_column in df.columns:
        return _split_grouped(df, labels, config, n_folds, random_state)
    else:
        return _split_simple(df, labels, config, n_folds, random_state)

Base classes

base

Abstract base class for dataset splitters and the SplitResult dataclass.

SplitResult dataclass

SplitResult(folds: dict[int, DataFrame], default_train_folds: list[int], default_val_folds: list[int], default_test_folds: list[int], stratify_column: str, group_column: str | None, split_metadata: dict = dict())

Output of any splitting operation.

train property

train: DataFrame

Default training set: all default train folds concatenated.

val property

val: DataFrame

Default validation set.

test property

test: DataFrame

Default test set.

get_fold

get_fold(fold_number: int) -> DataFrame

Get a single fold by number (1-indexed).

Source code in ecgbench/splitting/base.py
def get_fold(self, fold_number: int) -> pd.DataFrame:
    """Get a single fold by number (1-indexed)."""
    if fold_number not in self.folds:
        raise ValueError(
            f"Fold {fold_number} not found. Available: {sorted(self.folds)}"
        )
    return self.folds[fold_number]

get_kfold_split

get_kfold_split(val_fold: int, test_fold: int) -> tuple[DataFrame, DataFrame, DataFrame]

Get train/val/test for a custom k-fold rotation.

All folds except val_fold and test_fold become train.

Parameters:

Name Type Description Default
val_fold int

Fold number to use as validation.

required
test_fold int

Fold number to use as test.

required

Returns:

Type Description
tuple[DataFrame, DataFrame, DataFrame]

Tuple of (train_df, val_df, test_df).

Source code in ecgbench/splitting/base.py
def get_kfold_split(
    self, val_fold: int, test_fold: int
) -> tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame]:
    """Get train/val/test for a custom k-fold rotation.

    All folds except val_fold and test_fold become train.

    Args:
        val_fold: Fold number to use as validation.
        test_fold: Fold number to use as test.

    Returns:
        Tuple of (train_df, val_df, test_df).
    """
    train_folds = [f for f in self.folds if f != val_fold and f != test_fold]
    train = pd.concat([self.folds[f] for f in sorted(train_folds)], ignore_index=True)
    val = self.folds[val_fold]
    test = self.folds[test_fold]
    return train, val, test

DatasetSplitter

Bases: ABC

Abstract base for dataset-specific splitting logic.

Subclass this when a dataset has unusual structure that can't be handled by config alone (e.g., PTB-XL's SCP code -> superclass mapping). For simple datasets, use GenericSplitter which reads everything from config.

load_metadata abstractmethod

load_metadata(data_path: Path, config: DatasetConfig) -> DataFrame

Load the dataset's metadata CSV and normalise it.

Must return a DataFrame containing at minimum
  • config.record_id_column
  • config.signal_path_columns values
  • config.label_column (or derived stratification column)
  • config.patient_id_column (if applicable)
Source code in ecgbench/splitting/base.py
@abstractmethod
def load_metadata(self, data_path: Path, config: DatasetConfig) -> pd.DataFrame:
    """Load the dataset's metadata CSV and normalise it.

    Must return a DataFrame containing at minimum:
      - config.record_id_column
      - config.signal_path_columns values
      - config.label_column (or derived stratification column)
      - config.patient_id_column (if applicable)
    """
    ...

get_stratification_labels abstractmethod

get_stratification_labels(df: DataFrame, config: DatasetConfig) -> Series

Return a Series of categorical labels for stratification.

Must be aligned with df's index.

Source code in ecgbench/splitting/base.py
@abstractmethod
def get_stratification_labels(
    self, df: pd.DataFrame, config: DatasetConfig
) -> pd.Series:
    """Return a Series of categorical labels for stratification.

    Must be aligned with df's index.
    """
    ...

Registry

Splitters register themselves with @register("<config-slug>"). A strategy module must also be imported in ecgbench/splitting/strategies/__init__.py, otherwise the decorator never runs and lookup silently falls back to GenericSplitter.

registry

Splitter registry with automatic GenericSplitter fallback.

register

register(slug: str)

Decorator to register a splitter class for a dataset slug.

Source code in ecgbench/splitting/registry.py
def register(slug: str):
    """Decorator to register a splitter class for a dataset slug."""

    def wrapper(cls: type[DatasetSplitter]):
        _REGISTRY[slug] = cls
        return cls

    return wrapper

get_splitter

get_splitter(dataset_slug: str) -> DatasetSplitter

Get the splitter for a dataset. Falls back to GenericSplitter.

Source code in ecgbench/splitting/registry.py
def get_splitter(dataset_slug: str) -> DatasetSplitter:
    """Get the splitter for a dataset. Falls back to GenericSplitter."""
    from ecgbench.splitting.strategies.generic import GenericSplitter

    cls = _REGISTRY.get(dataset_slug, GenericSplitter)
    return cls()

Export

Fold CSVs carry minimal columns only — record ID, patient ID, signal paths, fold, default_split, plus is_valid/quality_issues in original/. Full metadata stays in the dataset's own CSV.

export

Fold CSV export for both original/ and clean/ versions.

Exported CSVs contain only the minimum columns needed to identify records in the original dataset: record ID, patient ID (if available), signal file paths, fold number, and default split assignment. Users join these back to the original dataset CSV to get full metadata (age, sex, labels, etc.).

export_splits

export_splits(split_result: SplitResult, validation_result: ValidationResult, output_dir: Path, config: DatasetConfig) -> dict

Export fold CSVs in both original/ and clean/ versions.

Exported CSVs contain only identification columns (record ID, patient ID, signal file paths) plus fold/split assignment. Full metadata stays in the original dataset CSV — users join on record_id when they need it.

Creates

output_dir/ original/ folds.csv train/fold_1.csv ... fold_N.csv val/fold_M.csv test/fold_K.csv clean/ folds.csv train/fold_1.csv ... fold_N.csv val/fold_M.csv test/fold_K.csv validation_report.json

Parameters:

Name Type Description Default
split_result SplitResult

SplitResult from split_dataset()

required
validation_result ValidationResult

ValidationResult from validate_dataset()

required
output_dir Path

Root output directory

required
config DatasetConfig

DatasetConfig

required

Returns:

Type Description
dict

dict with statistics: counts per split, per version

Source code in ecgbench/splitting/export.py
def export_splits(
    split_result: SplitResult,
    validation_result: ValidationResult,
    output_dir: Path,
    config: DatasetConfig,
) -> dict:
    """Export fold CSVs in both original/ and clean/ versions.

    Exported CSVs contain only identification columns (record ID, patient ID,
    signal file paths) plus fold/split assignment. Full metadata stays in the
    original dataset CSV — users join on record_id when they need it.

    Creates:
      output_dir/
        original/
          folds.csv
          train/fold_1.csv ... fold_N.csv
          val/fold_M.csv
          test/fold_K.csv
        clean/
          folds.csv
          train/fold_1.csv ... fold_N.csv
          val/fold_M.csv
          test/fold_K.csv
        validation_report.json

    Args:
        split_result: SplitResult from split_dataset()
        validation_result: ValidationResult from validate_dataset()
        output_dir: Root output directory
        config: DatasetConfig

    Returns:
        dict with statistics: counts per split, per version
    """
    output_dir = Path(output_dir)
    original_dir = output_dir / "original"
    clean_dir = output_dir / "clean"

    # Build the full master DataFrame by concatenating all folds with fold/split columns
    fold_to_split = _build_split_column(split_result)
    all_parts = []
    for fold_num, fold_df in sorted(split_result.folds.items()):
        part = fold_df.copy()
        part["fold"] = fold_num
        part["default_split"] = fold_to_split.get(fold_num, "train")
        all_parts.append(part)

    master_df = pd.concat(all_parts, ignore_index=True)

    # Merge validation info
    val_df = validation_result.original_df[
        [config.record_id_column, "is_valid", "quality_issues"]
    ].copy()
    val_df[config.record_id_column] = val_df[config.record_id_column].astype(
        master_df[config.record_id_column].dtype
    )

    master_df = master_df.merge(val_df, on=config.record_id_column, how="left")
    master_df["is_valid"] = master_df["is_valid"].fillna(True)
    master_df["quality_issues"] = master_df["quality_issues"].fillna("")

    # Sort by record_id for deterministic output
    master_df = master_df.sort_values(config.record_id_column).reset_index(drop=True)

    _check_zero_padded_identifiers(master_df, config)

    # --- Original version (minimal columns + quality flags) ---
    original_cols = _minimal_columns(config, include_quality=True)
    original_slim = _select_columns(master_df, original_cols)
    original_dir.mkdir(parents=True, exist_ok=True)
    original_slim.to_csv(original_dir / "folds.csv", index=False)
    _write_split_csvs(original_slim, original_dir, split_result, config)

    # --- Clean version (minimal columns, no quality flags) ---
    clean_rows = master_df[master_df["is_valid"]]
    clean_cols = _minimal_columns(config, include_quality=False)
    clean_slim = _select_columns(clean_rows, clean_cols).reset_index(drop=True)
    clean_dir.mkdir(parents=True, exist_ok=True)
    clean_slim.to_csv(clean_dir / "folds.csv", index=False)
    _write_split_csvs(clean_slim, clean_dir, split_result, config)

    # --- Validation report ---
    from ecgbench.validation.report import save_report

    save_report(validation_result, config, output_dir / "validation_report.json")

    # --- Statistics ---
    stats = {
        "original": {
            "total": len(original_slim),
            "train": int((original_slim["default_split"] == "train").sum()),
            "val": int((original_slim["default_split"] == "val").sum()),
            "test": int((original_slim["default_split"] == "test").sum()),
        },
        "clean": {
            "total": len(clean_slim),
            "train": int((clean_slim["default_split"] == "train").sum()),
            "val": int((clean_slim["default_split"] == "val").sum()),
            "test": int((clean_slim["default_split"] == "test").sum()),
        },
        "excluded": validation_result.excluded_records,
    }

    logger.info(
        "Export complete: original=%d, clean=%d, excluded=%d",
        stats["original"]["total"],
        stats["clean"]["total"],
        stats["excluded"],
    )

    return stats