Add disease history ablation modes
This commit is contained in:
57
tests/test_dataset_checkup.py
Normal file
57
tests/test_dataset_checkup.py
Normal file
@@ -0,0 +1,57 @@
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
from dataset import _ExpoBaseDataset
|
||||
from targets import CHECKUP_IDX
|
||||
from train_util import load_extra_info_types_file
|
||||
|
||||
|
||||
class CheckupSelectionTests(unittest.TestCase):
|
||||
@staticmethod
|
||||
def _base(extra_info_types):
|
||||
dataset = _ExpoBaseDataset.__new__(_ExpoBaseDataset)
|
||||
dataset.extra_info_types = list(extra_info_types)
|
||||
dataset.event_data = np.asarray(
|
||||
[
|
||||
[101, 10, CHECKUP_IDX],
|
||||
[101, 20, 2],
|
||||
[101, 30, 3],
|
||||
],
|
||||
dtype=np.float64,
|
||||
)
|
||||
return dataset
|
||||
|
||||
def test_explicit_empty_extra_info_removes_checkup(self):
|
||||
project_root = Path(__file__).resolve().parents[1]
|
||||
selected_types = load_extra_info_types_file(
|
||||
str(project_root / "extra_info_types_none.txt")
|
||||
)
|
||||
self.assertEqual(selected_types, [])
|
||||
dataset = self._base(selected_types)
|
||||
|
||||
rows = list(dataset._iter_patient_events(impute_no_event_gaps=False))
|
||||
|
||||
self.assertEqual(len(rows), 1)
|
||||
eid, times, labels = rows[0]
|
||||
self.assertEqual(eid, 101)
|
||||
np.testing.assert_array_equal(times, np.asarray([20, 30], dtype=np.float32))
|
||||
self.assertNotIn(CHECKUP_IDX, labels.tolist())
|
||||
|
||||
def test_selected_extra_info_keeps_checkup(self):
|
||||
dataset = self._base([11])
|
||||
|
||||
rows = list(dataset._iter_patient_events(impute_no_event_gaps=False))
|
||||
|
||||
self.assertEqual(len(rows), 1)
|
||||
_, times, labels = rows[0]
|
||||
np.testing.assert_array_equal(
|
||||
times,
|
||||
np.asarray([10, 20, 30], dtype=np.float32),
|
||||
)
|
||||
self.assertEqual(labels[0], CHECKUP_IDX)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
381
tests/test_disease_history_ablation.py
Normal file
381
tests/test_disease_history_ablation.py
Normal file
@@ -0,0 +1,381 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user