LuisKazuto23's picture
Deploy assistive robot study app
b38f323
Raw
History Blame Contribute Delete
1.61 kB
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