Train ATI-HO and finalize project outputs
This commit is contained in:
@@ -0,0 +1,166 @@
|
||||
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,
|
||||
}
|
||||
Reference in New Issue
Block a user