import unittest from pathlib import Path from unittest.mock import patch import numpy as np import torch from dataset import ( AllFutureHealthDataset, all_future_collate_fn, transform_disease_history, transform_disease_history_batch_at_position, ) from eval_data import validate_training_mode_config from losses import build_loss from models import DeepHealth from train_all_future import parse_args class DiseaseHistoryTransformTests(unittest.TestCase): def test_timed_ordered_and_set_representations(self): events = np.asarray([9, 4, 7], dtype=np.int64) times = np.asarray([50.0, 60.0, 65.0], dtype=np.float32) timed_events, timed_times, timed_query = transform_disease_history( events, times, 70.0, "timed" ) np.testing.assert_array_equal(timed_events, events) np.testing.assert_array_equal(timed_times, times) self.assertEqual(float(timed_query), 70.0) ordered_events, ordered_times, ordered_query = transform_disease_history( events, times, 70.0, "ordered" ) np.testing.assert_array_equal(ordered_events, events) np.testing.assert_array_equal( ordered_times, np.asarray([0.0, 1.0, 2.0], dtype=np.float32), ) self.assertEqual(float(ordered_query), 3.0) set_events, set_times, set_query = transform_disease_history( events, times, 70.0, "set" ) np.testing.assert_array_equal( set_events, np.asarray([4, 7, 9], dtype=np.int64), ) np.testing.assert_array_equal( set_times, np.zeros(3, dtype=np.float32), ) self.assertEqual(float(set_query), 0.0) def test_ordered_removes_calendar_time_but_keeps_order(self): events = np.asarray([9, 4, 7], dtype=np.int64) first = transform_disease_history( events, np.asarray([20.0, 21.0, 70.0], dtype=np.float32), 75.0, "ordered", ) second = transform_disease_history( events, np.asarray([50.0, 60.0, 65.0], dtype=np.float32), 70.0, "ordered", ) np.testing.assert_array_equal(first[0], second[0]) np.testing.assert_array_equal(first[1], second[1]) self.assertEqual(float(first[2]), float(second[2])) reversed_events = transform_disease_history( events[::-1], np.asarray([50.0, 60.0, 65.0], dtype=np.float32), 70.0, "ordered", )[0] self.assertFalse(np.array_equal(first[0], reversed_events)) def test_ordered_keeps_same_day_diseases_in_one_order_group(self): events, model_times, model_query = transform_disease_history( np.asarray([9, 4, 7], dtype=np.int64), np.asarray([50.0, 50.0, 65.0], dtype=np.float32), 70.0, "ordered", ) np.testing.assert_array_equal( events, np.asarray([9, 4, 7], dtype=np.int64), ) np.testing.assert_array_equal( model_times, np.asarray([0.0, 0.0, 1.0], dtype=np.float32), ) self.assertEqual(float(model_query), 2.0) def test_set_removes_order_and_calendar_time(self): first = transform_disease_history( np.asarray([9, 4, 7], dtype=np.int64), np.asarray([50.0, 60.0, 65.0], dtype=np.float32), 70.0, "set", ) second = transform_disease_history( np.asarray([7, 9, 4], dtype=np.int64), np.asarray([20.0, 21.0, 70.0], dtype=np.float32), 75.0, "set", ) np.testing.assert_array_equal(first[0], second[0]) np.testing.assert_array_equal(first[1], second[1]) self.assertEqual(float(first[2]), float(second[2])) def test_all_future_targets_stay_on_actual_time(self): patient = { "times": np.asarray([50.0, 60.0, 65.0, 75.0], dtype=np.float32), "labels": np.asarray([9, 4, 7, 12], dtype=np.int64), "t_obs": 75.0, "sex": 0, "other_type": np.zeros(0, dtype=np.int64), "other_value": np.zeros(0, dtype=np.float32), "other_value_kind": np.zeros(0, dtype=np.int64), "other_time": np.zeros(0, dtype=np.float32), } items = {} for mode in ("timed", "ordered", "set"): dataset = AllFutureHealthDataset.__new__(AllFutureHealthDataset) dataset.disease_history_mode = mode items[mode] = dataset._build_item(patient, 70.0) for item in items.values(): torch.testing.assert_close( item["future_targets"], torch.tensor([12], dtype=torch.long), ) torch.testing.assert_close( item["future_dt"], torch.tensor([5.0]), ) torch.testing.assert_close( item["exposure"], torch.tensor(5.0), ) torch.testing.assert_close( items["timed"]["time_seq"], torch.tensor([50.0, 60.0, 65.0]), ) self.assertEqual(float(items["timed"]["t_query"]), 70.0) torch.testing.assert_close( items["ordered"]["event_seq"], torch.tensor([9, 4, 7]), ) torch.testing.assert_close( items["ordered"]["time_seq"], torch.tensor([0.0, 1.0, 2.0]), ) self.assertEqual(float(items["ordered"]["t_query"]), 3.0) torch.testing.assert_close( items["set"]["event_seq"], torch.tensor([4, 7, 9]), ) torch.testing.assert_close( items["set"]["time_seq"], torch.zeros(3), ) self.assertEqual(float(items["set"]["t_query"]), 0.0) def test_batch_prefix_transform_masks_future_events(self): events = torch.tensor( [ [9, 4, 7, 12], [8, 5, 11, 0], ], dtype=torch.long, ) actual_times = torch.tensor( [ [50.0, 60.0, 65.0, 75.0], [45.0, 55.0, 80.0, 0.0], ] ) mask = events > 0 timed = transform_disease_history_batch_at_position( events, actual_times, mask, 1, "timed", vocab_size=20 ) torch.testing.assert_close(timed[0], events) torch.testing.assert_close(timed[1], actual_times) torch.testing.assert_close(timed[2], mask) torch.testing.assert_close(timed[3], torch.tensor([60.0, 55.0])) ordered = transform_disease_history_batch_at_position( events, actual_times, mask, 1, "ordered", vocab_size=20 ) torch.testing.assert_close( ordered[2], torch.tensor( [ [True, True, False, False], [True, True, False, False], ] ), ) torch.testing.assert_close(ordered[3], torch.tensor([2.0, 2.0])) disease_set = transform_disease_history_batch_at_position( events, actual_times, mask, 1, "set", vocab_size=20 ) torch.testing.assert_close( disease_set[0], torch.tensor( [ [4, 9, 0, 0], [5, 8, 0, 0], ] ), ) torch.testing.assert_close( disease_set[2], torch.tensor( [ [True, True, False, False], [True, True, False, False], ] ), ) torch.testing.assert_close(disease_set[3], torch.zeros(2)) class DiseaseSetModelTests(unittest.TestCase): @staticmethod def _model(): return DeepHealth( vocab_size=16, n_embd=24, n_head=4, n_layer=2, n_types=1, n_cont_types=0, n_categories=1, cont_type_ids=[], n_bins=4, target_mode="all_future", time_mode="relative", dist_mode="weibull", dropout=0.0, model_architecture="traj_mixer_v5", ) def test_equal_time_query_is_permutation_invariant(self): torch.manual_seed(7) model = self._model().eval() event_seq = torch.tensor( [ [4, 7, 9], [9, 4, 7], ], dtype=torch.long, ) time_seq = torch.zeros(2, 3) empty_long = torch.zeros(2, 0, dtype=torch.long) empty_float = torch.zeros(2, 0) with torch.inference_mode(): hidden = model( event_seq=event_seq, time_seq=time_seq, sex=torch.zeros(2, dtype=torch.long), padding_mask=torch.ones(2, 3, dtype=torch.bool), t_query=torch.zeros(2), other_type=empty_long, other_value=empty_float, other_value_kind=empty_long, other_time=empty_float, ) torch.testing.assert_close(hidden[0], hidden[1], atol=1e-6, rtol=1e-6) def test_ordered_and_set_support_finite_weibull_loss(self): patient = { "times": np.asarray([50.0, 60.0, 65.0, 75.0], dtype=np.float32), "labels": np.asarray([9, 4, 7, 12], dtype=np.int64), "t_obs": 75.0, "sex": 0, "other_type": np.zeros(0, dtype=np.int64), "other_value": np.zeros(0, dtype=np.float32), "other_value_kind": np.zeros(0, dtype=np.int64), "other_time": np.zeros(0, dtype=np.float32), } criterion = build_loss("weibull", ignored_idx={0, 1}) for mode in ("ordered", "set"): dataset = AllFutureHealthDataset.__new__(AllFutureHealthDataset) dataset.disease_history_mode = mode item = dataset._build_item(patient, 70.0) batch = all_future_collate_fn([item, item]) model = self._model() hidden = model( event_seq=batch["event_seq"], time_seq=batch["time_seq"], sex=batch["sex"], padding_mask=batch["padding_mask"], t_query=batch["t_query"], other_type=batch["other_type"], other_value=batch["other_value"], other_value_kind=batch["other_value_kind"], other_time=batch["other_time"], ) loss = criterion( logits=model.calc_risk(hidden), weibull_rho=model.calc_weibull_rho(hidden), targets=batch["future_targets"], dt=batch["future_dt"], exposure=batch["exposure"], ) self.assertTrue(torch.isfinite(loss), msg=f"{mode} loss={loss}") class DiseaseHistoryConfigTests(unittest.TestCase): def test_training_cli_accepts_ordered_ablation(self): project_root = Path(__file__).resolve().parents[1] with patch( "sys.argv", [ "train_all_future.py", "--disease_history_mode", "ordered", "--model_architecture", "traj_mixer_v5", "--time_mode", "relative", "--dist_mode", "weibull", "--extra_info_types_file", str(project_root / "extra_info_types_none.txt"), ], ): args = parse_args() self.assertEqual(args.disease_history_mode, "ordered") self.assertEqual(args.extra_info_types, []) def test_legacy_config_defaults_to_timed(self): validate_training_mode_config( { "model_target_mode": "all_future", "time_mode": "relative", "dist_mode": "weibull", "model_architecture": "traj_mixer_v5", "extra_info_types": [11, 66, 67], } ) def test_ordered_config_requires_exact_ablation_setup(self): validate_training_mode_config( { "model_target_mode": "all_future", "time_mode": "relative", "dist_mode": "weibull", "model_architecture": "traj_mixer_v5", "extra_info_types": [], "disease_history_mode": "ordered", } ) with self.assertRaises(ValueError): validate_training_mode_config( { "model_target_mode": "all_future", "time_mode": "relative", "dist_mode": "weibull", "model_architecture": "traj_mixer_v5", "extra_info_types": [11], "disease_history_mode": "ordered", } ) if __name__ == "__main__": unittest.main()