From 7aadd32e141abb36f38443b0dcb9bedc381dff4b Mon Sep 17 00:00:00 2001 From: Daryl Okeke Date: Wed, 19 Aug 2026 21:54:28 -0500 Subject: [PATCH] Fix sign toggle in removal-based interpretability metrics original_class_probs aliased y_probs, which is computed once before the loop over percentages. Negating NEGATIVE-class entries in place therefore flipped the sign on every iteration instead of applying it per percentage, so a sample's score depended on where its percentage sat in the list. Clone before negating. Same for ablated_probs, whose in-place negation also corrupted the debug output that prints it as P(class=1). --- pyhealth/metrics/interpretability/base.py | 4 ++-- tests/core/test_interp_metrics.py | 27 +++++++++++++++++++++++ 2 files changed, 29 insertions(+), 2 deletions(-) diff --git a/pyhealth/metrics/interpretability/base.py b/pyhealth/metrics/interpretability/base.py index ef388402b..46dd234ee 100644 --- a/pyhealth/metrics/interpretability/base.py +++ b/pyhealth/metrics/interpretability/base.py @@ -453,10 +453,10 @@ def compute( ) # Compute probability drop - original_class_probs = y_probs + original_class_probs = y_probs.clone() original_class_probs[neg_mask] = -original_class_probs[neg_mask] - ablated_class_probs = ablated_probs + ablated_class_probs = ablated_probs.clone() ablated_class_probs[neg_mask] = -ablated_class_probs[neg_mask] prob_drop = torch.zeros(batch_size, device=y_probs.device) diff --git a/tests/core/test_interp_metrics.py b/tests/core/test_interp_metrics.py index c415f1a9c..e2d1a1376 100644 --- a/tests/core/test_interp_metrics.py +++ b/tests/core/test_interp_metrics.py @@ -16,6 +16,7 @@ from pyhealth.metrics.interpretability import ( ComprehensivenessMetric, Evaluator, + SampleClass, SufficiencyMetric, threshold_sample_filter, ) @@ -449,6 +450,32 @@ def test_percentage_sensitivity(self): self.assertTrue(torch.isfinite(torch.tensor(score_10))) self.assertTrue(torch.isfinite(torch.tensor(score_50))) + def test_negative_class_scores_independent_of_percentage_order(self): + """Test that a negative-class sample's score at a percentage is order-independent.""" + attributions = self._create_attributions(self.batch) + + def negative_filter(y_probs, classifier_type): + return torch.full( + (y_probs.shape[0],), + SampleClass.NEGATIVE, + dtype=torch.long, + device=y_probs.device, + ) + + def score_at_20(percentages): + comp = ComprehensivenessMetric( + self.model, + percentages=percentages, + ablation_strategy="zero", + sample_filter=negative_filter, + ) + detailed = comp.compute( + self.batch, attributions, return_per_percentage=True + ) + return detailed[20] + + torch.testing.assert_close(score_at_20([20]), score_at_20([10, 20])) + def test_attribution_shape_mismatch(self): """Test that mismatched attribution shapes are handled gracefully.""" # Skip this test - shape mismatches may not always raise errors