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()