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, }