diff --git a/docs/api/processors/pyhealth.processors.NestedSequenceProcessor.rst b/docs/api/processors/pyhealth.processors.NestedSequenceProcessor.rst index cce30d08c..afda7eb59 100644 --- a/docs/api/processors/pyhealth.processors.NestedSequenceProcessor.rst +++ b/docs/api/processors/pyhealth.processors.NestedSequenceProcessor.rst @@ -7,6 +7,23 @@ Handles nested sequences like drug recommendation history where each sample contains a list of visits, and each visit contains a list of codes. For example: [["code1", "code2"], ["code3"], ["code4", "code5", "code6"]] +Output width +------------ + +Each visit becomes a row of the same width: the longest visit seen in ``fit()`` +plus ``padding``. Shorter visits are padded with ````. Longer visits, which +can occur when processors are fitted on the training split and applied to +validation or test data, are truncated to that width, keeping the first codes, +with a one-time warning. Set ``padding`` to leave room for longer visits. +``NestedFloatsProcessor``, ``DeepNestedSequenceProcessor`` and +``DeepNestedFloatsProcessor`` truncate codes or values per visit the same way. + +In PyHealth 2.0.2 and earlier, longer visits were not truncated. Samples then had +different widths and batching them failed with a tensor shape mismatch (the deep +processors failed already while processing the sample). + +See ``examples/nested_sequence_fit_on_train.py``. + .. autoclass:: pyhealth.processors.NestedSequenceProcessor :members: :undoc-members: diff --git a/examples/nested_sequence_fit_on_train.py b/examples/nested_sequence_fit_on_train.py new file mode 100644 index 000000000..11778c249 --- /dev/null +++ b/examples/nested_sequence_fit_on_train.py @@ -0,0 +1,44 @@ +"""Fit a nested-sequence processor on training data, then apply it to test data. + +Fitting processors on the training split only avoids leaking test information, +but a test patient can then have a visit with more codes than any training +visit. ``NestedSequenceProcessor`` keeps the width it learned in ``fit()``: +longer visits are truncated (keeping the first codes, with a one-time warning) +so every sample has the same shape and batches stack. Use ``padding`` to leave +room for longer visits. + +Runs in a second on CPU with synthetic data; no download. + +Usage: + python examples/nested_sequence_fit_on_train.py +""" + +import torch + +from pyhealth.processors import NestedSequenceProcessor + +train = [ + {"visits": [["I10", "E11"], ["I10"]]}, + {"visits": [["J45", "E11", "I10"]]}, # longest training visit: 3 codes +] +test_visits = [["I10"], ["J45", "E11", "I10", "N18", "E78"]] # 5 codes + + +def main(): + processor = NestedSequenceProcessor() + processor.fit(train, "visits") + print("fitted width:", processor.size()) + + rows = processor.process(test_visits) # warns once: 5 codes > 3 + print("test sample shape:", tuple(rows.shape)) + print("batch of train + test stacks:", tuple(torch.cat( + [processor.process(train[1]["visits"]), rows]).shape)) + + roomy = NestedSequenceProcessor(padding=2) # width = 3 + 2 + roomy.fit(train, "visits") + print("with padding=2, width:", roomy.size(), + "-> test shape:", tuple(roomy.process(test_visits).shape)) + + +if __name__ == "__main__": + main() diff --git a/pyhealth/processors/base_processor.py b/pyhealth/processors/base_processor.py index 06bbe0c7c..6f8f5ed63 100644 --- a/pyhealth/processors/base_processor.py +++ b/pyhealth/processors/base_processor.py @@ -1,9 +1,25 @@ from abc import ABC, abstractmethod from enum import Enum +import logging from typing import Any, Dict, List, Iterable import torch +logger = logging.getLogger(__name__) + + +def _warn_truncated(processor: Any, what: str, got: int, limit: int) -> None: + """Logs once per processor that an input was cut to the fitted width.""" + if getattr(processor, "_warned_truncation", False): + return + processor._warned_truncation = True + logger.warning( + "%s: an input has %d %s, more than the %d seen in fit(); keeping the " + "first %d. Inputs longer than the fitted data are truncated so every " + "sample has the same shape. Pass a larger `padding` to keep more.", + type(processor).__name__, got, what, limit, limit, + ) + class ModalityType(str, Enum): """Standard modality identifiers for routing in UnifiedMultimodalEmbeddingModel. diff --git a/pyhealth/processors/deep_nested_sequence_processor.py b/pyhealth/processors/deep_nested_sequence_processor.py index 24683f54b..6911453a4 100644 --- a/pyhealth/processors/deep_nested_sequence_processor.py +++ b/pyhealth/processors/deep_nested_sequence_processor.py @@ -3,7 +3,7 @@ import torch from . import register_processor -from .base_processor import FeatureProcessor, TokenProcessorInterface +from .base_processor import FeatureProcessor, TokenProcessorInterface, _warn_truncated @register_processor("deep_nested_sequence") @@ -159,7 +159,10 @@ def process(self, value: List[List[List[Any]]]) -> torch.Tensor: else: indices.append(self.code_vocab[code]) - # Pad codes dimension to max_inner_len + # Truncate to, or pad to, max_inner_len from fit() + if len(indices) > self._max_inner_len: + _warn_truncated(self, "codes in a visit", len(indices), self._max_inner_len) + indices = indices[: self._max_inner_len] while len(indices) < self._max_inner_len: indices.append(pad_token) @@ -341,7 +344,10 @@ def process(self, value: List[List[List[float]]]) -> torch.Tensor: else: values.append(0.0) - # Pad inner dimension + # Truncate to, or pad to, max_inner_len from fit() + if len(values) > self._max_inner_len: + _warn_truncated(self, "values in a visit", len(values), self._max_inner_len) + values = values[: self._max_inner_len] while len(values) < self._max_inner_len: if self.forward_fill: values.append(float("nan")) diff --git a/pyhealth/processors/nested_sequence_processor.py b/pyhealth/processors/nested_sequence_processor.py index 060b5e3f8..2518329f4 100644 --- a/pyhealth/processors/nested_sequence_processor.py +++ b/pyhealth/processors/nested_sequence_processor.py @@ -3,7 +3,7 @@ import torch from . import register_processor -from .base_processor import FeatureProcessor, TokenProcessorInterface +from .base_processor import FeatureProcessor, TokenProcessorInterface, _warn_truncated @register_processor("nested_sequence") @@ -18,7 +18,8 @@ class NestedSequenceProcessor(FeatureProcessor, TokenProcessorInterface): The processor: 1. Builds a vocabulary from all codes across all samples 2. Encodes codes to indices - 3. Pads inner sequences to the maximum sequence length found during fit + 3. Pads inner sequences to the maximum sequence length found during fit, + and truncates longer ones (keeping the first codes) 4. Returns a 2D tensor of shape (num_visits, max_codes_per_visit) Special tokens: @@ -28,8 +29,9 @@ class NestedSequenceProcessor(FeatureProcessor, TokenProcessorInterface): Args: padding: Additional padding to add on top of the observed maximum inner sequence length. The actual padding length will be observed_max + padding. - This ensures the processor can handle sequences longer than those in the - training data. Default: 0 (no extra padding). + Inner sequences longer than this are truncated to it (keeping the first + elements, with a one-time warning) so every sample has the same width; + raise padding to keep longer sequences. Default: 0 (no extra padding). Examples: >>> processor = NestedSequenceProcessor() @@ -142,7 +144,10 @@ def process(self, value: List[List[Any]]) -> torch.Tensor: else: indices.append(self.code_vocab[code]) - # Pad to maximum inner length + # Truncate to, or pad to, the maximum inner length from fit() + if len(indices) > self._max_inner_len: + _warn_truncated(self, "codes in a visit", len(indices), self._max_inner_len) + indices = indices[: self._max_inner_len] while len(indices) < self._max_inner_len: indices.append(pad_token) @@ -229,8 +234,9 @@ class NestedFloatsProcessor(FeatureProcessor): Default is True. padding: Additional padding to add on top of the observed maximum inner sequence length. The actual padding length will be observed_max + padding. - This ensures the processor can handle sequences longer than those in the - training data. Default: 0 (no extra padding). + Inner sequences longer than this are truncated to it (keeping the first + elements, with a one-time warning) so every sample has the same width; + raise padding to keep longer sequences. Default: 0 (no extra padding). Examples: >>> processor = NestedFloatsProcessor() @@ -331,7 +337,10 @@ def process(self, value: List[List[float]]) -> torch.Tensor: else: values.append(0.0) - # Pad to maximum inner length + # Truncate to, or pad to, the maximum inner length from fit() + if len(values) > self._max_inner_len: + _warn_truncated(self, "values in a visit", len(values), self._max_inner_len) + values = values[: self._max_inner_len] while len(values) < self._max_inner_len: if self.forward_fill: values.append(float("nan")) diff --git a/tests/core/test_nested_processor_truncation.py b/tests/core/test_nested_processor_truncation.py new file mode 100644 index 000000000..115633b31 --- /dev/null +++ b/tests/core/test_nested_processor_truncation.py @@ -0,0 +1,85 @@ +"""Nested processors must return their fitted width for longer inputs. + +When processors are fitted on a training split, a validation or test sample can +have a longer visit (or more visits per group) than anything seen in fit(). +The output must still have the fitted width, otherwise samples have ragged +shapes and batching fails. +""" + +import unittest + +import torch + +from pyhealth.processors import ( + DeepNestedFloatsProcessor, + DeepNestedSequenceProcessor, + NestedFloatsProcessor, + NestedSequenceProcessor, +) + +LOGGER = DEEP_LOGGER = "pyhealth.processors.base_processor" + + +class TestNestedSequenceTruncation(unittest.TestCase): + def setUp(self): + self.p = NestedSequenceProcessor() + self.p.fit([{"v": [["a", "b"], ["c"]]}], "v") # fitted width 2 + + def test_longer_visit_is_truncated_to_fitted_width(self): + with self.assertLogs(LOGGER, "WARNING"): + out = self.p.process([["a", "b", "c", "d"]]) + self.assertEqual(tuple(out.shape), (1, 2)) + self.assertEqual(out.tolist(), [[self.p.code_vocab["a"], self.p.code_vocab["b"]]]) + + def test_batch_of_longer_and_shorter_visits_stacks(self): + with self.assertLogs(LOGGER, "WARNING"): + rows = [self.p.process([["a"]]), self.p.process([["a", "b", "c"]])] + self.assertEqual(tuple(torch.cat(rows).shape), (2, 2)) + + def test_warns_once_per_processor(self): + with self.assertLogs(LOGGER, "WARNING") as logs: + self.p.process([["a", "b", "c"]]) + self.p.process([["a", "b", "c", "d"]]) + self.assertEqual(len(logs.records), 1) + + def test_padding_keeps_longer_visits(self): + p = NestedSequenceProcessor(padding=2) + p.fit([{"v": [["a", "b"]]}], "v") + with self.assertNoLogs(LOGGER, "WARNING"): + out = p.process([["a", "b", "c", "d"]]) + self.assertEqual(tuple(out.shape), (1, 4)) + + +class TestNestedFloatsTruncation(unittest.TestCase): + def test_longer_visit_is_truncated(self): + for forward_fill in (True, False): + with self.subTest(forward_fill=forward_fill): + p = NestedFloatsProcessor(forward_fill=forward_fill) + p.fit([{"v": [[1.0, 2.0]]}], "v") + with self.assertLogs(LOGGER, "WARNING"): + out = p.process([[1.0, 2.0, 3.0, 4.0]]) + self.assertEqual(tuple(out.shape), (1, 2)) + self.assertEqual(out.tolist(), [[1.0, 2.0]]) + + +class TestDeepNestedTruncation(unittest.TestCase): + def test_codes_per_visit_are_truncated(self): + p = DeepNestedSequenceProcessor() + p.fit([{"v": [[["a", "b"], ["c"]]]}], "v") # 2 visits x 2 codes + with self.assertLogs(DEEP_LOGGER, "WARNING"): + out = p.process([[["a", "b", "c"], ["a"]]]) # a visit with 3 codes + self.assertEqual(tuple(out.shape), (1, 2, 2)) + + def test_float_values_per_visit_are_truncated(self): + for forward_fill in (True, False): + with self.subTest(forward_fill=forward_fill): + p = DeepNestedFloatsProcessor(forward_fill=forward_fill) + p.fit([{"v": [[[1.0, 2.0], [3.0]]]}], "v") + with self.assertLogs(DEEP_LOGGER, "WARNING"): + out = p.process([[[1.0, 2.0, 9.0], [3.0]]]) + self.assertEqual(tuple(out.shape), (1, 2, 2)) + self.assertEqual(out[0, 0].tolist(), [1.0, 2.0]) + + +if __name__ == "__main__": + unittest.main()