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