Add Q3 MoFE router visualizations and explanations
This commit is contained in:
@@ -0,0 +1,52 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user