Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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 ``<pad>``. 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:
Expand Down
44 changes: 44 additions & 0 deletions examples/nested_sequence_fit_on_train.py
Original file line number Diff line number Diff line change
@@ -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()
16 changes: 16 additions & 0 deletions pyhealth/processors/base_processor.py
Original file line number Diff line number Diff line change
@@ -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.
Expand Down
12 changes: 9 additions & 3 deletions pyhealth/processors/deep_nested_sequence_processor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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"))
Expand Down
25 changes: 17 additions & 8 deletions pyhealth/processors/nested_sequence_processor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand All @@ -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:
Expand All @@ -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()
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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"))
Expand Down
85 changes: 85 additions & 0 deletions tests/core/test_nested_processor_truncation.py
Original file line number Diff line number Diff line change
@@ -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()
Loading