53 lines
2.1 KiB
Python
53 lines
2.1 KiB
Python
from __future__ import annotations
|
|
|
|
import itertools
|
|
import unittest
|
|
|
|
import numpy as np
|
|
|
|
from .run_experiments import COALITIONS, exact_pair_interactions, exact_shapley, _scale_features
|
|
|
|
|
|
class ExactAttributionTests(unittest.TestCase):
|
|
def test_additive_game_shapley_and_interaction(self) -> None:
|
|
weights = (1.25, -0.5, 2.0)
|
|
values = {
|
|
coalition: 3.0 + sum(weights[player] for player in coalition)
|
|
for coalition in COALITIONS
|
|
}
|
|
np.testing.assert_allclose(exact_shapley(values), weights, atol=1e-12)
|
|
for interaction in exact_pair_interactions(values).values():
|
|
self.assertAlmostEqual(interaction, 0.0, places=12)
|
|
|
|
def test_pair_interaction_is_reported_with_standard_half_weight(self) -> None:
|
|
values = {}
|
|
for coalition in COALITIONS:
|
|
value = float(len(coalition))
|
|
if 0 in coalition and 1 in coalition:
|
|
value += 2.0
|
|
values[coalition] = value
|
|
interactions = exact_pair_interactions(values)
|
|
self.assertAlmostEqual(interactions[(0, 1)], 1.0, places=12)
|
|
self.assertAlmostEqual(interactions[(0, 2)], 0.0, places=12)
|
|
self.assertAlmostEqual(interactions[(1, 2)], 0.0, places=12)
|
|
|
|
def test_robust_scaling_supports_single_cases_and_validation_batches(self) -> None:
|
|
features = (np.full((2, 4, 2), 3.0, np.float32),)
|
|
centers = (np.asarray([1.0, 1.0], np.float32),)
|
|
scales = (np.asarray([2.0, 2.0], np.float32),)
|
|
single_mask = np.ones((4, 1), dtype=bool)
|
|
single = _scale_features((features[0][0],), single_mask, centers, scales)[0]
|
|
self.assertEqual(single.shape, (4, 2))
|
|
np.testing.assert_allclose(single, 1.0)
|
|
batch_mask = np.ones((2, 4, 1), dtype=bool)
|
|
batch_mask[1, 2:, 0] = False
|
|
batch = _scale_features(features, batch_mask, centers, scales)[0]
|
|
self.assertEqual(batch.shape, (2, 4, 2))
|
|
np.testing.assert_allclose(batch[0], 1.0)
|
|
np.testing.assert_allclose(batch[1, :2], 1.0)
|
|
np.testing.assert_allclose(batch[1, 2:], 0.0)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|