Spaces:
Sleeping
Sleeping
| from __future__ import annotations | |
| import random | |
| from collections import defaultdict | |
| from typing import Callable, TypeVar | |
| T = TypeVar("T") | |
| def stratified_split( | |
| items: list[T], | |
| *, | |
| label_getter: Callable[[T], str], | |
| train_ratio: float = 0.8, | |
| val_ratio: float = 0.1, | |
| seed: int = 42, | |
| ) -> tuple[list[T], list[T], list[T]]: | |
| if not 0 < train_ratio < 1: | |
| raise ValueError("train_ratio must be between 0 and 1") | |
| if not 0 <= val_ratio < 1: | |
| raise ValueError("val_ratio must be between 0 and 1") | |
| if train_ratio + val_ratio >= 1: | |
| raise ValueError("train_ratio + val_ratio must be < 1") | |
| grouped: dict[str, list[T]] = defaultdict(list) | |
| for item in items: | |
| grouped[label_getter(item)].append(item) | |
| rng = random.Random(seed) | |
| train: list[T] = [] | |
| val: list[T] = [] | |
| test: list[T] = [] | |
| for label_items in grouped.values(): | |
| rows = list(label_items) | |
| rng.shuffle(rows) | |
| n = len(rows) | |
| n_train = int(round(n * train_ratio)) | |
| n_val = int(round(n * val_ratio)) | |
| if n_train + n_val > n: | |
| n_val = max(0, n - n_train) | |
| n_test = n - n_train - n_val | |
| if n >= 3 and n_test == 0: | |
| if n_train > 1: | |
| n_train -= 1 | |
| elif n_val > 1: | |
| n_val -= 1 | |
| n_test = n - n_train - n_val | |
| train.extend(rows[:n_train]) | |
| val.extend(rows[n_train : n_train + n_val]) | |
| test.extend(rows[n_train + n_val :]) | |
| rng.shuffle(train) | |
| rng.shuffle(val) | |
| rng.shuffle(test) | |
| return train, val, test | |