diff --git a/pyhealth/models/micron.py b/pyhealth/models/micron.py index ffee2612c..4087cacda 100644 --- a/pyhealth/models/micron.py +++ b/pyhealth/models/micron.py @@ -77,6 +77,8 @@ def compute_reconstruction_loss( Returns: torch.tensor: Mean squared reconstruction loss value. """ + if logits.size(1) < 2: + return logits.new_zeros(()) rec_loss = torch.mean( torch.square( torch.sigmoid(logits[:, 1:, :]) @@ -246,21 +248,20 @@ def _ensure_tensor(self, feature_key: str, value) -> torch.Tensor: return torch.tensor(value, dtype=torch.long) return torch.tensor(value, dtype=torch.float) - def _pool_embedding(self, x: torch.Tensor) -> torch.Tensor: - if x.dim() == 4: + def _pool_embedding(self, feature_key: str, x: torch.Tensor) -> torch.Tensor: + if isinstance(self.feature_processors[feature_key], SequenceProcessor): + # Flat code lists describe one visit, not successive visits. + x = x.sum(dim=1, keepdim=True) + elif x.dim() == 4: x = x.sum(dim=2) if x.dim() == 2: x = x.unsqueeze(1) - # Make sure temporal dimension (dim=1) matches the longest sequence - if x.size(1) == 1: - # Repeat to handle shorter sequences - x = x.repeat(1, 2, 1) return x def _create_mask(self, feature_key: str, value: torch.Tensor) -> torch.Tensor: processor = self.feature_processors[feature_key] if isinstance(processor, SequenceProcessor): - mask = value != 0 + mask = (value != 0).any(dim=1, keepdim=True) elif isinstance(processor, StageNetProcessor): if value.dim() >= 3: mask = torch.any(value != 0, dim=-1) @@ -332,7 +333,7 @@ def forward(self, **kwargs) -> Dict[str, torch.Tensor]: for feature_key in self.feature_keys: x = embedded[feature_key] mask = masks[feature_key] - x = self._pool_embedding(x) + x = self._pool_embedding(feature_key, x) patient_emb.append(x) # Concatenate along last dim: [batch, seq_len, embedding_dim * n_features] diff --git a/tests/core/test_micron.py b/tests/core/test_micron.py index 59107f907..0552a50d3 100644 --- a/tests/core/test_micron.py +++ b/tests/core/test_micron.py @@ -124,9 +124,8 @@ def test_model_with_embedding(self): self.assertIn("embed", ret) self.assertEqual(ret["embed"].shape[0], 2) # batch size - expected_seq_len = max(len(data_batch["conditions"][0]), len(data_batch["procedures"][0])) expected_feature_dim = self.model.embedding_dim * len(self.model.feature_keys) - self.assertEqual(ret["embed"].shape[1], expected_seq_len) + self.assertEqual(ret["embed"].shape[1], 1) self.assertEqual(ret["embed"].shape[2], expected_feature_dim) def test_custom_hyperparameters(self): @@ -151,6 +150,57 @@ def test_custom_hyperparameters(self): self.assertIn("loss", ret) self.assertIn("y_prob", ret) + def test_features_with_unequal_sequence_lengths(self): + """Unequal code lists must both contribute to the visit's prediction.""" + samples = [{ + "conditions": ["a", "b", "c"], + "procedures": ["x"], + "drugs": ["d1", "d2"], + }] + dataset = create_sample_dataset( + samples=samples, + input_schema={"conditions": "sequence", "procedures": "sequence"}, + output_schema={"drugs": "multilabel"}, + ) + model = MICRON(dataset=dataset, embedding_dim=2, hidden_dim=2) + # Positive, small weights avoid cancellation and sigmoid saturation. + with torch.no_grad(): + for parameter in model.parameters(): + parameter.fill_(0.1) + for layer in model.embedding_model.embedding_layers.values(): + layer.weight[0].zero_() + data_batch = next(iter(get_dataloader(dataset, batch_size=1, shuffle=False))) + ret = model(**data_batch, embed=True) + self.assertTrue(torch.isfinite(ret["loss"])) + # Each flat list belongs to the same single visit. + self.assertEqual(ret["embed"].shape[1], 1) + + ret["y_prob"].sum().backward() + for field, codes in {"conditions": ["a", "b", "c"], "procedures": ["x"]}.items(): + vocab = dataset.input_processors[field].code_vocab + grad = model.embedding_model.embedding_layers[field].weight.grad + for code in codes: + self.assertTrue(torch.isfinite(grad[vocab[code]]).all()) + self.assertGreater(grad[vocab[code]].abs().sum().item(), 0) + + def test_nested_sequences_preserve_visits(self): + samples = [{ + "conditions": [["a", "b"], ["c"]], + "procedures": [["x"], ["y", "z"]], + "drugs": ["d1"], + }] + dataset = create_sample_dataset( + samples=samples, + input_schema={"conditions": "nested_sequence", "procedures": "nested_sequence"}, + output_schema={"drugs": "multilabel"}, + ) + model = MICRON(dataset=dataset) + batch = next(iter(get_dataloader(dataset, batch_size=1, shuffle=False))) + result = model(**batch, embed=True) + self.assertEqual(result["embed"].shape[1], 2) + self.assertTrue(torch.isfinite(result["loss"])) + result["loss"].backward() + def test_ddi_adjacency_matrix(self): """Test the drug-drug interaction adjacency matrix generation.""" ddi_matrix = self.model.generate_ddi_adj()