Spaces:
Sleeping
Sleeping
File size: 1,605 Bytes
b38f323 | 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 | 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
|