Files
DeepHealth/tests/test_disease_history_ablation.py

382 lines
13 KiB
Python
Raw Normal View History

2026-07-29 13:53:15 +08:00
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()