Files

167 lines
6.6 KiB
Python

from __future__ import annotations
import itertools
import math
from typing import Any, Sequence
import numpy as np
import torch
from torch import nn
def ensemble_forward(
models: Sequence[nn.Module],
xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
masks: torch.Tensor,
*,
details: bool = True,
) -> dict[str, Any]:
"""Average additive parameters first, then decode the ensemble prediction."""
outputs = []
for model in models:
try:
outputs.append(model(xs, masks, return_details=details))
except TypeError:
outputs.append(model(xs, masks))
if "params" not in outputs[0]:
logits = torch.stack([output["logits"] for output in outputs], dim=0).mean(dim=0)
probabilities = torch.stack(
[torch.softmax(output["logits"], dim=-1) for output in outputs], dim=0
).mean(dim=0)
intensity = torch.stack([output["intensity"] for output in outputs], dim=0).mean(dim=0)
result = {
"logits": logits,
"probabilities": probabilities,
"predicted_class": logits.argmax(dim=-1),
"intensity": intensity,
}
if "utility" in outputs[0]:
result["utility"] = torch.stack(
[output["utility"] for output in outputs], dim=0
).mean(dim=0)
return result
result: dict[str, Any] = {}
averaged = ("params", "baseline", "main_effects", "pair_effects", "mask_logits")
for key in averaged:
if key in outputs[0]:
result[key] = torch.stack([output[key] for output in outputs], dim=0).mean(dim=0)
params = result["params"]
logits = params[:, :3]
probabilities = torch.softmax(logits, dim=-1)
nu_negative = 3.0 * torch.sigmoid(params[:, 3])
nu_positive = 3.0 * torch.sigmoid(params[:, 4])
predicted_class = logits.argmax(dim=-1)
intensity = torch.where(
predicted_class == 0,
-nu_negative,
torch.where(predicted_class == 2, nu_positive, torch.zeros_like(nu_positive)),
)
result.update(
{
"logits": logits,
"probabilities": probabilities,
"predicted_class": predicted_class,
"intensity": intensity,
"nu_negative": nu_negative,
"nu_positive": nu_positive,
"soft_intensity": probabilities[:, 2] * nu_positive - probabilities[:, 0] * nu_negative,
}
)
return result
def shapley_from_eight(values: np.ndarray) -> np.ndarray:
"""Exact three-player Shapley values from the eight coalition values."""
values = np.asarray(values, dtype=np.float64)
if values.shape[-1] != 8:
raise ValueError("the three-modality game requires exactly eight coalition values")
result = np.zeros((*values.shape[:-1], 3), dtype=np.float64)
factorial = math.factorial
for modality in range(3):
others = [i for i in range(3) if i != modality]
for size in range(3):
weight = factorial(size) * factorial(2 - size) / factorial(3)
for subset in itertools.combinations(others, size):
before = sum(1 << item for item in subset)
after = before | (1 << modality)
result[..., modality] += weight * (values[..., after] - values[..., before])
return result
def analytic_class_shapley(details: dict[str, torch.Tensor], target: torch.Tensor, other: torch.Tensor) -> np.ndarray:
"""Closed form for the fixed target-vs-runner-up logit margin."""
main = details["main_effects"]
pairs = details["pair_effects"]
delta = torch.zeros((main.shape[0], 3), dtype=main.dtype, device=main.device)
rows = torch.arange(main.shape[0], device=main.device)
delta[rows, target] = 1.0
delta[rows, other] = -1.0
contributions = main.clone()
pair_modalities = ((0, 1), (0, 2), (1, 2))
for pair_idx, (left, right) in enumerate(pair_modalities):
contributions[:, left] = contributions[:, left] + 0.5 * pairs[:, pair_idx]
contributions[:, right] = contributions[:, right] + 0.5 * pairs[:, pair_idx]
values = torch.einsum("bi,bmi->bm", delta, contributions[..., :3])
return values.detach().cpu().numpy().astype(np.float64)
@torch.inference_mode()
def exact_shapley_audit(
models: Sequence[nn.Module],
xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
masks: torch.Tensor,
*,
batch_size: int = 128,
) -> dict[str, np.ndarray]:
"""Compare analytic output-parameter Shapley with exact 8-coalition values.
The class target and runner-up are fixed from each sample's full-input
prediction. A second exact game is computed for the decoded hard intensity;
that value is nonlinear and is not compared with the analytic formula.
"""
n = masks.shape[0]
full = ensemble_forward(models, xs, masks, details=True)
logits = full["logits"]
target = logits.argmax(dim=-1)
ranked = logits.argsort(dim=-1, descending=True)
other = ranked[:, 1]
analytic = analytic_class_shapley(full, target, other)
margin_values = np.zeros((n, 8), dtype=np.float64)
intensity_values = np.zeros((n, 8), dtype=np.float64)
for coalition in range(8):
for start in range(0, n, batch_size):
end = min(n, start + batch_size)
current_mask = masks[start:end].clone()
for modality in range(3):
if not coalition & (1 << modality):
current_mask[..., modality] = False
current_xs = tuple(x[start:end] for x in xs)
output = ensemble_forward(models, current_xs, current_mask, details=False)
local_rows = torch.arange(end - start, device=logits.device)
local_target = target[start:end]
local_other = other[start:end]
margin = (
output["logits"][local_rows, local_target]
- output["logits"][local_rows, local_other]
)
margin_values[start:end, coalition] = margin.detach().cpu().numpy()
intensity_values[start:end, coalition] = output["intensity"].detach().cpu().numpy()
exact = shapley_from_eight(margin_values)
exact_intensity = shapley_from_eight(intensity_values)
error = np.abs(analytic - exact)
tolerance = 1e-6 + 1e-5 * np.abs(exact)
return {
"analytic_class": analytic,
"exact_class": exact,
"class_abs_error": error,
"class_pass": error <= tolerance,
"exact_intensity": exact_intensity,
"coalition_margin": margin_values,
"coalition_intensity": intensity_values,
"target_class": target.detach().cpu().numpy(),
"runner_up_class": other.detach().cpu().numpy(),
"full_output": full,
}