diff --git a/pyhealth/models/transformer_deid.py b/pyhealth/models/transformer_deid.py index 964942f5a..6cdfc38ef 100644 --- a/pyhealth/models/transformer_deid.py +++ b/pyhealth/models/transformer_deid.py @@ -243,7 +243,12 @@ def deidentify(self, text: str, redact: str = "[REDACTED]") -> str: dummy_labels = " ".join(["O"] * len(words)) self.eval() with torch.no_grad(): - result = self(text=[text], labels=[dummy_labels]) + result = self( + **{ + self.feature_key: [text], + self.label_key: [dummy_labels], + } + ) preds = result["logit"][0].argmax(dim=-1) y_true = result["y_true"][0] diff --git a/tests/core/test_transformer_deid.py b/tests/core/test_transformer_deid.py index a796663c1..5c49ddb61 100644 --- a/tests/core/test_transformer_deid.py +++ b/tests/core/test_transformer_deid.py @@ -195,5 +195,31 @@ def test_deidentify_custom_redact_marker(self): ) +class TestTransformerDeIDCustomFieldNames(unittest.TestCase): + class StubModel(TransformerDeID): + def __init__(self): + torch.nn.Module.__init__(self) + self.feature_key = "clinical_note" + self.label_key = "bio_tags" + + def forward(self, **kwargs): + texts = kwargs[self.feature_key] + labels = kwargs[self.label_key] + word_count = len(texts[0].split()) + self.received_labels = labels + return { + "logit": torch.zeros(1, word_count, len(LABEL_VOCAB)), + "y_true": torch.zeros(1, word_count, dtype=torch.long), + } + + def test_deidentify_uses_configured_field_names(self): + model = self.StubModel() + + result = model.deidentify("Patient Jane Doe") + + self.assertEqual(result, "Patient Jane Doe") + self.assertEqual(model.received_labels, ["O O O"]) + + if __name__ == "__main__": unittest.main()