From c5a3633e84f1a296c7e362e7c1926aead4063359 Mon Sep 17 00:00:00 2001 From: Brendan McMahan Date: Wed, 26 Aug 2026 12:05:57 -0700 Subject: [PATCH] Fix sample weight generation in membership_inference_attack_test. * Replace `rng.randn` with `rng.rand` for sample weights in MIA test helpers to ensure non-negative sample weights required by `RandomForestClassifier`'s bootstrap sampling logic. PiperOrigin-RevId: 971422996 --- .../membership_inference_attack_test.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/tensorflow_privacy/privacy/privacy_tests/membership_inference_attack/membership_inference_attack_test.py b/tensorflow_privacy/privacy/privacy_tests/membership_inference_attack/membership_inference_attack_test.py index 4644814d..a0f84317 100644 --- a/tensorflow_privacy/privacy/privacy_tests/membership_inference_attack/membership_inference_attack_test.py +++ b/tensorflow_privacy/privacy/privacy_tests/membership_inference_attack/membership_inference_attack_test.py @@ -83,8 +83,9 @@ def get_multilabel_test_input_with_sample_weights(n_train, n_test): logits_test=rng.randn(n_test, num_classes) + 0.2, labels_train=get_multihot_labels_for_test(n_train, num_classes), labels_test=get_multihot_labels_for_test(n_test, num_classes), - sample_weight_train=rng.randn(n_train, 1), - sample_weight_test=rng.randn(n_test, 1)) + sample_weight_train=rng.rand(n_train, 1), + sample_weight_test=rng.rand(n_test, 1), + ) def get_test_input_logits_only(n_train, n_test): @@ -101,8 +102,9 @@ def get_test_input_logits_only_with_sample_weights(n_train, n_test): return AttackInputData( logits_train=rng.randn(n_train, 5) + 0.2, logits_test=rng.randn(n_test, 5) + 0.2, - sample_weight_train=rng.randn(n_train, 1), - sample_weight_test=rng.randn(n_test, 1)) + sample_weight_train=rng.rand(n_train, 1), + sample_weight_test=rng.rand(n_test, 1), + ) class MockTrainedAttacker(object):