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()