Files
DeepHealth/tests/test_continuous_value_scaling.py

143 lines
4.5 KiB
Python

import unittest
import numpy as np
import torch
import torch.nn as nn
from torch.utils.data import Subset
from models import OtherInfoTokenizer
from train_util import fit_continuous_robust_scaler
class _ToyAllFutureDataset:
def __init__(self):
self.cont_type_ids = [1, 3]
self.n_types = 4
self.patients = [
self._patient([1, 3], [0.0, 10.0]),
self._patient([1, 3], [1.0, 10.0]),
self._patient([1, 3], [2.0, 10.0]),
self._patient([1, 3], [3.0, 10.0]),
self._patient([1, 3], [4.0, 10.0]),
self._patient([1, 3], [1000.0, 999.0]),
]
@staticmethod
def _patient(types, values):
return {
"other_type": np.asarray(types, dtype=np.int64),
"other_value": np.asarray(values, dtype=np.float32),
"other_value_kind": np.ones(len(types), dtype=np.int64),
}
def __len__(self):
return len(self.patients)
class _CaptureContinuousEncoder(nn.Module):
def __init__(self, n_embd):
super().__init__()
self.n_embd = n_embd
self.last_type = None
self.last_value = None
def forward(self, cont_type_idx, value):
self.last_type = cont_type_idx.detach().clone()
self.last_value = value.detach().clone()
return value[:, None].expand(-1, self.n_embd)
class ContinuousValueScalingTests(unittest.TestCase):
def test_fit_uses_only_training_subset_and_handles_constant_features(self):
dataset = _ToyAllFutureDataset()
train_subset = Subset(dataset, np.asarray([0, 1, 2, 3, 4]))
stats = fit_continuous_robust_scaler(dataset, train_subset)
self.assertEqual(stats.cont_type_ids, (1, 3))
np.testing.assert_array_equal(stats.observation_count, np.asarray([5, 5]))
np.testing.assert_allclose(stats.center, np.asarray([2.0, 10.0]))
np.testing.assert_allclose(stats.scale, np.asarray([2.0, 1.0]))
def test_tokenizer_standardizes_only_continuous_values(self):
tokenizer = OtherInfoTokenizer(
n_embd=4,
n_types=4,
n_cont_types=2,
n_categories=3,
cont_type_ids=[1, 3],
continuous_value_scaling="robust",
continuous_value_center=[10.0, 100.0],
continuous_value_scale=[2.0, 20.0],
)
capture = _CaptureContinuousEncoder(n_embd=4)
tokenizer.cont_value_encoder = capture
tokenizer(
other_type=torch.tensor([[1, 2, 3]], dtype=torch.long),
other_value=torch.tensor([[14.0, 1.0, 80.0]]),
other_value_kind=torch.tensor([[1, 2, 1]], dtype=torch.long),
)
torch.testing.assert_close(capture.last_type, torch.tensor([0, 1]))
torch.testing.assert_close(capture.last_value, torch.tensor([2.0, -1.0]))
def test_scaler_buffers_round_trip_in_new_checkpoint(self):
tokenizer = OtherInfoTokenizer(
n_embd=4,
n_types=4,
n_cont_types=2,
n_categories=2,
cont_type_ids=[1, 3],
continuous_value_scaling="robust",
continuous_value_center=[2.0, 10.0],
continuous_value_scale=[1.5, 4.0],
)
state = tokenizer.state_dict()
self.assertIn("continuous_value_center", state)
self.assertIn("continuous_value_scale", state)
restored = OtherInfoTokenizer(
n_embd=4,
n_types=4,
n_cont_types=2,
n_categories=2,
cont_type_ids=[1, 3],
continuous_value_scaling="robust",
)
restored.load_state_dict(state, strict=True)
torch.testing.assert_close(
restored.continuous_value_center,
torch.tensor([2.0, 10.0]),
)
torch.testing.assert_close(
restored.continuous_value_scale,
torch.tensor([1.5, 4.0]),
)
def test_legacy_none_mode_keeps_old_state_dict_schema(self):
tokenizer = OtherInfoTokenizer(
n_embd=4,
n_types=4,
n_cont_types=2,
n_categories=2,
cont_type_ids=[1, 3],
)
state = tokenizer.state_dict()
self.assertNotIn("continuous_value_center", state)
self.assertNotIn("continuous_value_scale", state)
restored = OtherInfoTokenizer(
n_embd=4,
n_types=4,
n_cont_types=2,
n_categories=2,
cont_type_ids=[1, 3],
)
restored.load_state_dict(state, strict=True)
if __name__ == "__main__":
unittest.main()