Files
DeepHealth/tests/test_export_disease_state_xiad.py

172 lines
6.1 KiB
Python
Raw Normal View History

2026-09-01 10:59:09 +08:00
import tempfile
import unittest
from pathlib import Path
import numpy as np
from export_disease_state_xiad import (
build_prevalent_mask,
compute_xiad,
load_label_axis_rows,
validate_axis_mapping_csv,
validate_horizons,
validate_source_token_order,
weibull_probability,
write_axis_mapping_csv,
)
class _DummyDataset:
def __init__(self) -> None:
self.samples = [
{
"eid": 101,
"event_seq": np.asarray([3, 4], dtype=np.int64),
"time_seq": np.asarray([1.0, 2.0], dtype=np.float32),
"target_event_seq": np.asarray([4, 5], dtype=np.int64),
"target_time_seq": np.asarray([2.0, 4.0], dtype=np.float32),
},
{
"eid": 102,
"event_seq": np.asarray([3], dtype=np.int64),
"time_seq": np.asarray([1.5], dtype=np.float32),
"target_event_seq": np.asarray([6], dtype=np.int64),
"target_time_seq": np.asarray([5.0], dtype=np.float32),
},
]
class DiseaseStateXiadTests(unittest.TestCase):
def test_label_order_defines_disease_axis_and_excludes_death(self) -> None:
with tempfile.TemporaryDirectory() as tmp_dir:
labels_path = Path(tmp_dir) / "labels.csv"
labels_path.write_text(
"A00 (cholera)\nI10 Essential hypertension\nDeath\n",
encoding="utf-8",
)
label_rows = load_label_axis_rows(labels_path)
source = {
"tokens/column": np.asarray([0, 1, 2], dtype=np.int64),
"tokens/token_id": np.asarray([3, 4, 5], dtype=np.int64),
"tokens/label_code": np.asarray([b"A00", b"I10", b"Death"]),
"tokens/label_text": np.asarray(
[b"A00 (cholera)", b"I10 Essential hypertension", b"Death"]
),
"tokens/outcome_type": np.asarray(
[b"disease", b"disease", b"death"]
),
}
disease_rows = validate_source_token_order(source, label_rows)
self.assertEqual([row["column"] for row in disease_rows], [0, 1])
self.assertEqual([row["source_column"] for row in disease_rows], [0, 1])
self.assertEqual([row["label_index"] for row in disease_rows], [0, 1])
self.assertEqual([row["token_id"] for row in disease_rows], [3, 4])
self.assertEqual([row["code"] for row in disease_rows], ["A00", "I10"])
with tempfile.TemporaryDirectory() as tmp_dir:
mapping_path = Path(tmp_dir) / "icd10_columns.csv"
write_axis_mapping_csv(mapping_path, disease_rows)
validate_axis_mapping_csv(mapping_path, disease_rows)
header = mapping_path.read_text(encoding="utf-8").splitlines()[0]
self.assertEqual(
header,
"column,source_column,label_index,token_id,code,name,label_text",
)
def test_source_order_mismatch_is_rejected(self) -> None:
label_rows = [
{
"label_index": 0,
"token_id": 3,
"code": "A00",
"name": "cholera",
"label_text": "A00 (cholera)",
"outcome_type": "disease",
}
]
source = {
"tokens/column": np.asarray([0], dtype=np.int64),
"tokens/token_id": np.asarray([3], dtype=np.int64),
"tokens/label_code": np.asarray([b"A01"]),
"tokens/label_text": np.asarray([b"A01 wrong order"]),
"tokens/outcome_type": np.asarray([b"disease"]),
}
with self.assertRaises(ValueError):
validate_source_token_order(source, label_rows)
def test_weibull_probability_combines_shape_and_scale(self) -> None:
shape = np.asarray([[2.0, 1.0]], dtype=np.float32)
scale = np.asarray([[10.0, 4.0]], dtype=np.float32)
result = weibull_probability(
shape,
scale,
np.asarray([5.0, 10.0], dtype=np.float64),
)
expected = np.asarray(
[
[
[1.0 - np.exp(-0.25), 1.0 - np.exp(-1.0)],
[1.0 - np.exp(-1.25), 1.0 - np.exp(-2.5)],
]
],
dtype=np.float32,
)
np.testing.assert_allclose(result, expected, rtol=1e-6, atol=1e-7)
def test_prevalent_disease_state_is_one(self) -> None:
result = compute_xiad(
shape=np.asarray([[2.0, 1.0]], dtype=np.float32),
scale=np.asarray([[10.0, 4.0]], dtype=np.float32),
horizons=np.asarray([5.0], dtype=np.float64),
prevalent=np.asarray([[True, False]]),
)
self.assertEqual(float(result[0, 0, 0]), 1.0)
self.assertAlmostEqual(
float(result[0, 1, 0]),
1.0 - np.exp(-1.25),
places=6,
)
def test_invalid_nonprevalent_parameter_produces_nan(self) -> None:
result = compute_xiad(
shape=np.asarray([[1.0, 1.0]], dtype=np.float32),
scale=np.asarray([[np.nan, -1.0]], dtype=np.float32),
horizons=np.asarray([5.0], dtype=np.float64),
prevalent=np.asarray([[True, False]]),
)
self.assertEqual(float(result[0, 0, 0]), 1.0)
self.assertTrue(np.isnan(result[0, 1, 0]))
def test_prevalence_uses_disease_time_at_or_before_landmark(self) -> None:
result = build_prevalent_mask(
dataset=_DummyDataset(),
dataset_indices=np.asarray([0, 1], dtype=np.int64),
landmark_age=2.0,
token_to_column={3: 0, 4: 1, 5: 2, 6: 3},
n_diseases=4,
)
expected = np.asarray(
[
[True, True, False, False],
[True, False, False, False],
]
)
np.testing.assert_array_equal(result, expected)
def test_horizons_must_be_positive_and_unique(self) -> None:
with self.assertRaises(ValueError):
validate_horizons([0.0, 5.0])
with self.assertRaises(ValueError):
validate_horizons([5.0, 5.0])
if __name__ == "__main__":
unittest.main()