163 lines
5.7 KiB
Python
163 lines
5.7 KiB
Python
import argparse
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
from unittest.mock import patch
|
|
|
|
import numpy as np
|
|
|
|
from eval_data import split_indices
|
|
from evaluate_auc import make_eval_subset
|
|
from evaluate_auc_v2 import make_eval_indices
|
|
|
|
|
|
class _DummyDataset:
|
|
def __init__(self, eids: list[int]) -> None:
|
|
self.samples = [{"eid": eid} for eid in eids]
|
|
|
|
def __len__(self) -> int:
|
|
return len(self.samples)
|
|
|
|
def __getitem__(self, index: int) -> dict[str, int]:
|
|
return self.samples[index]
|
|
|
|
|
|
class AUCEidSplitTests(unittest.TestCase):
|
|
def test_both_auc_evaluators_default_to_ukb_test_eid_file(self) -> None:
|
|
dataset = _DummyDataset([101, 102, 103, 104, 105])
|
|
args = argparse.Namespace(
|
|
eval_split="test",
|
|
dataset_subset_size=None,
|
|
test_eid_file=None,
|
|
)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
root = Path(tmp_dir)
|
|
(root / "ukb_test_eid.csv").write_text(
|
|
"eid\n104\n102\n",
|
|
encoding="utf-8",
|
|
)
|
|
with patch.object(Path, "cwd", return_value=root):
|
|
subset, legacy_indices = make_eval_subset(dataset, args, {})
|
|
landmark_indices = make_eval_indices(dataset, args, {})
|
|
|
|
expected = np.asarray([1, 3], dtype=np.int64)
|
|
np.testing.assert_array_equal(legacy_indices, expected)
|
|
np.testing.assert_array_equal(landmark_indices, expected)
|
|
self.assertEqual(subset.indices, expected.tolist())
|
|
|
|
def test_subset_size_is_applied_after_eid_selection(self) -> None:
|
|
dataset = _DummyDataset([201, 202, 203, 204])
|
|
args = argparse.Namespace(
|
|
eval_split="test",
|
|
dataset_subset_size=1,
|
|
test_eid_file=None,
|
|
)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
eid_path = Path(tmp_dir) / "test.csv"
|
|
eid_path.write_text("eid\n202\n204\n", encoding="utf-8")
|
|
cfg = {"test_eid_file": str(eid_path)}
|
|
_, legacy_indices = make_eval_subset(dataset, args, cfg)
|
|
landmark_indices = make_eval_indices(dataset, args, cfg)
|
|
|
|
expected = np.asarray([1], dtype=np.int64)
|
|
np.testing.assert_array_equal(legacy_indices, expected)
|
|
np.testing.assert_array_equal(landmark_indices, expected)
|
|
|
|
def test_cli_test_eid_file_overrides_config(self) -> None:
|
|
dataset = _DummyDataset([301, 302, 303])
|
|
|
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
root = Path(tmp_dir)
|
|
config_path = root / "config.csv"
|
|
cli_path = root / "cli.csv"
|
|
config_path.write_text("eid\n301\n", encoding="utf-8")
|
|
cli_path.write_text("eid\n303\n", encoding="utf-8")
|
|
args = argparse.Namespace(
|
|
eval_split="test",
|
|
dataset_subset_size=None,
|
|
test_eid_file=str(cli_path),
|
|
)
|
|
cfg = {"test_eid_file": str(config_path)}
|
|
|
|
_, legacy_indices = make_eval_subset(dataset, args, cfg)
|
|
landmark_indices = make_eval_indices(dataset, args, cfg)
|
|
|
|
expected = np.asarray([2], dtype=np.int64)
|
|
np.testing.assert_array_equal(legacy_indices, expected)
|
|
np.testing.assert_array_equal(landmark_indices, expected)
|
|
|
|
def test_empty_test_eid_file_explicitly_uses_ratio_split(self) -> None:
|
|
dataset = _DummyDataset(list(range(20)))
|
|
args = argparse.Namespace(
|
|
eval_split="test",
|
|
dataset_subset_size=None,
|
|
test_eid_file="",
|
|
)
|
|
cfg = {
|
|
"train_ratio": 0.7,
|
|
"val_ratio": 0.15,
|
|
"test_ratio": 0.15,
|
|
"seed": 7,
|
|
}
|
|
|
|
expected = split_indices(20, 0.7, 0.15, 0.15, 7)[2]
|
|
_, legacy_indices = make_eval_subset(dataset, args, cfg)
|
|
landmark_indices = make_eval_indices(dataset, args, cfg)
|
|
|
|
np.testing.assert_array_equal(legacy_indices, expected)
|
|
np.testing.assert_array_equal(landmark_indices, expected)
|
|
|
|
def test_non_test_split_does_not_read_test_eid_file(self) -> None:
|
|
dataset = _DummyDataset(list(range(20)))
|
|
args = argparse.Namespace(
|
|
eval_split="val",
|
|
dataset_subset_size=None,
|
|
test_eid_file="missing.csv",
|
|
)
|
|
cfg = {
|
|
"train_ratio": 0.7,
|
|
"val_ratio": 0.15,
|
|
"test_ratio": 0.15,
|
|
"seed": 11,
|
|
}
|
|
|
|
expected = split_indices(20, 0.7, 0.15, 0.15, 11)[1]
|
|
_, legacy_indices = make_eval_subset(dataset, args, cfg)
|
|
landmark_indices = make_eval_indices(dataset, args, cfg)
|
|
|
|
np.testing.assert_array_equal(legacy_indices, expected)
|
|
np.testing.assert_array_equal(landmark_indices, expected)
|
|
|
|
def test_missing_or_nonmatching_eid_file_fails_closed(self) -> None:
|
|
dataset = _DummyDataset([401, 402])
|
|
args = argparse.Namespace(
|
|
eval_split="test",
|
|
dataset_subset_size=None,
|
|
test_eid_file=None,
|
|
)
|
|
|
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
root = Path(tmp_dir)
|
|
missing = root / "missing.csv"
|
|
with self.assertRaisesRegex(FileNotFoundError, "EID split file"):
|
|
make_eval_indices(
|
|
dataset,
|
|
args,
|
|
{"test_eid_file": str(missing)},
|
|
)
|
|
|
|
nonmatching = root / "nonmatching.csv"
|
|
nonmatching.write_text("eid\n999\n", encoding="utf-8")
|
|
with self.assertRaisesRegex(ValueError, "No dataset patients"):
|
|
make_eval_indices(
|
|
dataset,
|
|
args,
|
|
{"test_eid_file": str(nonmatching)},
|
|
)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|