Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -179,3 +179,4 @@ cython_debug/

# numpy data
*.npy
!tests/eva/assets/vision/datasets/cc_ccii/CC-CCII_public/data/*
2 changes: 1 addition & 1 deletion .python-version
Original file line number Diff line number Diff line change
@@ -1 +1 @@
3.10.14
3.12.8
127 changes: 121 additions & 6 deletions pdm.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

2 changes: 2 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,7 @@ vision = [
"scipy>=1.14.0",
"monai>=1.3.2",
"einops>=0.8.1",
"medmnist>=3.0.2",
]
language = [
"datasets<4.0.0,>=2.19.0",
Expand Down Expand Up @@ -102,6 +103,7 @@ all = [
"litellm>=1.61.8",
"backoff>=2.2.1",
"json-repair>=0.52.0",
"medmnist>=3.0.2",
]

[project.scripts]
Expand Down
2 changes: 1 addition & 1 deletion src/eva/core/data/samplers/classification/balanced.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
""""Balanced class sampler for data loading."""
"""Balanced class sampler for data loading."""

from loguru import logger
from typing_extensions import override
Expand Down
47 changes: 40 additions & 7 deletions src/eva/core/data/splitting/stratified.py
Original file line number Diff line number Diff line change
@@ -1,16 +1,17 @@
"""Functions for stratified splitting."""

from typing import Any, List, Sequence, Tuple
from typing import Any, Iterable, List, Sequence, Tuple

import numpy as np


def stratified_split(
samples: Sequence[Any],
targets: Sequence[Any],
samples: Sequence[Any] | Iterable[Any],
targets: Sequence[Any] | Iterable[Any],
train_ratio: float,
val_ratio: float,
test_ratio: float = 0.0,
groups: Sequence[Any] | Iterable[Any] | None = None,
seed: int = 42,
) -> Tuple[List[int], List[int], List[int] | None]:
"""Splits the samples into stratified train, validation, and test (optional) sets.
Expand All @@ -21,19 +22,51 @@ def stratified_split(
train_ratio: The ratio of the training set.
val_ratio: The ratio of the validation set.
test_ratio: The ratio of the test set (optional).
groups: Optional group labels for group-wise stratification.
seed: The seed for reproducibility.

Returns:
The indices of the train, validation, and test sets.
"""
if len(samples) != len(targets):
samples_seq = samples if isinstance(samples, (list, tuple)) else list(samples)
targets_seq = targets if isinstance(targets, (list, tuple)) else list(targets)

if train_ratio + val_ratio + test_ratio > 1.0:
raise ValueError("The sum of the ratios must be lower or equal to 1")

if len(samples_seq) != len(targets_seq):
raise ValueError("The number of samples and targets must be equal.")
if train_ratio + val_ratio + (test_ratio or 0) > 1.0:
raise ValueError("The sum of the ratios must be lower or equal to 1.")

if groups is not None:
groups_seq = groups if isinstance(groups, (list, tuple)) else list(groups)
if len(groups_seq) != len(samples_seq):
raise ValueError("The number of samples and groups must be equal.")

unique_groups, group_indices = np.unique(groups_seq, return_inverse=True)
group_targets = np.array(
[targets_seq[np.where(group_indices == i)[0][0]] for i in range(len(unique_groups))]
)

train_g, val_g, test_g = stratified_split(
Comment thread
nkaenzig marked this conversation as resolved.
samples=unique_groups.tolist(),
targets=group_targets.tolist(),
train_ratio=train_ratio,
val_ratio=val_ratio,
test_ratio=test_ratio,
seed=seed,
)

def map_indices(g_list):
if g_list is None:
return []
selected_groups = unique_groups[g_list]
return np.where(np.isin(groups_seq, selected_groups))[0].tolist()

return map_indices(train_g), map_indices(val_g), map_indices(test_g) or None

use_all_samples = train_ratio + val_ratio + test_ratio == 1
random_generator = np.random.default_rng(seed)
unique_classes, y_indices = np.unique(targets, return_inverse=True)
unique_classes, y_indices = np.unique(targets_seq, return_inverse=True)
n_classes = unique_classes.shape[0]

train_indices, val_indices, test_indices = [], [], []
Expand Down
8 changes: 8 additions & 0 deletions src/eva/vision/data/datasets/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,12 +3,16 @@
from eva.vision.data.datasets.classification import (
BACH,
BRACS,
CC_CCII,
CRC,
LUNA25,
MHIST,
PANDA,
BreaKHis,
Camelyon16,
GleasonArvaniti,
NoduleMNIST3D,
OrganMNIST3D,
PANDASmall,
PatchCamelyon,
UniToPatho,
Expand Down Expand Up @@ -36,11 +40,15 @@
"BRACS",
"BTCV",
"Camelyon16",
"CC_CCII",
"CoNSeP",
"CRC",
"LUNA25",
"EmbeddingsSegmentationDataset",
"FLARE22",
"GleasonArvaniti",
"NoduleMNIST3D",
"OrganMNIST3D",
"KiTS23",
"LiTS17",
"MHIST",
Expand Down
8 changes: 8 additions & 0 deletions src/eva/vision/data/datasets/classification/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,13 @@
from eva.vision.data.datasets.classification.bracs import BRACS
from eva.vision.data.datasets.classification.breakhis import BreaKHis
from eva.vision.data.datasets.classification.camelyon16 import Camelyon16
from eva.vision.data.datasets.classification.cc_ccii import CC_CCII
from eva.vision.data.datasets.classification.crc import CRC
from eva.vision.data.datasets.classification.gleason_arvaniti import GleasonArvaniti
from eva.vision.data.datasets.classification.luna25 import LUNA25
from eva.vision.data.datasets.classification.mhist import MHIST
from eva.vision.data.datasets.classification.nodule_mnist_3d import NoduleMNIST3D
from eva.vision.data.datasets.classification.organ_mnist_3d import OrganMNIST3D
from eva.vision.data.datasets.classification.panda import PANDA, PANDASmall
from eva.vision.data.datasets.classification.patch_camelyon import PatchCamelyon
from eva.vision.data.datasets.classification.unitopatho import UniToPatho
Expand All @@ -17,9 +21,13 @@
"BreaKHis",
"BRACS",
"Camelyon16",
"CC_CCII",
"CRC",
"GleasonArvaniti",
"LUNA25",
"MHIST",
"NoduleMNIST3D",
"OrganMNIST3D",
"PatchCamelyon",
"UniToPatho",
"WsiClassificationDataset",
Expand Down
Loading
Loading