382 lines
13 KiB
Python
382 lines
13 KiB
Python
|
|
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()
|