172 lines
6.1 KiB
Python
172 lines
6.1 KiB
Python
|
|
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()
|