from __future__ import annotations from typing import Any, Sequence import numpy as np import torch from torch import nn from .attribution import ensemble_forward def _margin_values( models: Sequence[nn.Module], xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor], masks: np.ndarray, target: int, other: int, *, batch_size: int = 64, ) -> np.ndarray: device = xs[0].device values: list[np.ndarray] = [] with torch.inference_mode(): for start in range(0, len(masks), batch_size): end = min(len(masks), start + batch_size) current_mask = torch.as_tensor(masks[start:end], dtype=torch.bool, device=device) repeated = tuple(x.expand(end - start, -1, -1).contiguous() for x in xs) output = ensemble_forward(models, repeated, current_mask, details=False) margin = output["logits"][:, target] - output["logits"][:, other] values.append(margin.detach().cpu().numpy().astype(np.float64)) return np.concatenate(values) if values else np.empty(0, dtype=np.float64) def _top_stability(previous: np.ndarray, current: np.ndarray, k: int = 5) -> float: old = set(np.argsort(-np.abs(previous), kind="stable")[:k].tolist()) new = set(np.argsort(-np.abs(current), kind="stable")[:k].tolist()) return float(len(old & new) / max(1, len(old | new))) @torch.inference_mode() def hierarchical_owen_one( models: Sequence[nn.Module], xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor], mask: torch.Tensor, *, seed: int, bins_per_modality: int = 10, start_permutations: int = 8, max_permutations: int = 64, batch_size: int = 64, ) -> dict[str, Any]: """Estimate a three-group Owen allocation over relative-progress segments. The outer permutation orders modalities. Each modality's 10 temporal bins are then added in a random inner permutation. This is a sampled Owen value, not a perturbation of arbitrary individual feature dimensions. """ model_output = ensemble_forward(models, xs, mask, details=False) logits = model_output["logits"][0] target = int(logits.argmax().item()) other = int(logits.argsort(descending=True)[1].item()) full_margin = float((logits[target] - logits[other]).item()) no_features = torch.zeros_like(mask, dtype=torch.bool) baseline = ensemble_forward(models, xs, no_features, details=False)["logits"][0] base_margin = float((baseline[target] - baseline[other]).item()) base_mask = mask[0].detach().cpu().numpy().astype(bool) steps = base_mask.shape[0] boundaries = np.linspace(0, steps, bins_per_modality + 1).round().astype(int) slices = [(int(boundaries[k]), int(boundaries[k + 1])) for k in range(bins_per_modality)] rng = np.random.default_rng(seed) draws: list[np.ndarray] = [] previous_mean: np.ndarray | None = None stability = float("nan") schedule = [start_permutations] while schedule[-1] < max_permutations: schedule.append(min(max_permutations, schedule[-1] * 2)) next_target = schedule[0] stopping_status = "max_permutations" while len(draws) < max_permutations: needed = min(8, max_permutations - len(draws)) all_masks: list[np.ndarray] = [] player_order: list[list[tuple[int, int]]] = [] for _ in range(needed): current = np.zeros_like(base_mask, dtype=bool) order: list[tuple[int, int]] = [] outer = rng.permutation(3) for modality in outer: for bin_index in rng.permutation(bins_per_modality): left, right = slices[int(bin_index)] current[left:right, int(modality)] = base_mask[left:right, int(modality)] order.append((int(modality), int(bin_index))) all_masks.append(current.copy()) player_order.append(order) scores = _margin_values( models, xs, np.stack(all_masks), target, other, batch_size=batch_size, ) cursor = 0 for order in player_order: previous_score = base_margin draw = np.zeros((3, bins_per_modality), dtype=np.float64) for modality, bin_index in order: current_score = float(scores[cursor]) cursor += 1 draw[modality, bin_index] = current_score - previous_score previous_score = current_score draws.append(draw) if len(draws) >= next_target: current_mean = np.mean(np.stack(draws), axis=0) flat_mean = current_mean.reshape(-1) standard_error = np.std(np.stack(draws), axis=0, ddof=1) / np.sqrt(len(draws)) se_mean = float(np.mean(standard_error)) if previous_mean is not None: stability = _top_stability(previous_mean, flat_mean) se_limit = max(0.02, 0.10 * abs(full_margin - base_margin)) if stability >= 0.8 and se_mean <= se_limit: stopping_status = "stable" break previous_mean = flat_mean next_idx = next((i for i, value in enumerate(schedule) if value > len(draws)), None) if next_idx is None: break next_target = schedule[next_idx] draw_array = np.stack(draws) mean = draw_array.mean(axis=0) standard_error = draw_array.std(axis=0, ddof=1) / np.sqrt(len(draws)) if len(draws) > 1 else np.full_like(mean, np.nan) conservation = float(mean.sum() - (full_margin - base_margin)) return { "target_class": target, "runner_up_class": other, "full_margin": full_margin, "baseline_margin": base_margin, "contribution": mean, "standard_error": standard_error, "permutations": len(draws), "stopping_status": stopping_status, "top5_jaccard_last_check": stability, "local_conservation_residual": conservation, "bin_slices": slices, } @torch.inference_mode() def fidelity_audit_one( models: Sequence[nn.Module], xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor], mask: torch.Tensor, contribution: np.ndarray, *, sample_id: str, seed: int, bins_per_modality: int = 10, random_replicates: int = 20, ) -> list[dict[str, Any]]: """Deletion/retention at 10/20/30%, matched by modality and segment count.""" base_mask = mask[0].detach().cpu().numpy().astype(bool) steps = base_mask.shape[0] boundaries = np.linspace(0, steps, bins_per_modality + 1).round().astype(int) slices = [(int(boundaries[k]), int(boundaries[k + 1])) for k in range(bins_per_modality)] full = ensemble_forward(models, xs, mask, details=False) ranked = full["logits"][0].argsort(descending=True) target, other = int(ranked[0].item()), int(ranked[1].item()) full_margin = float((full["logits"][0, target] - full["logits"][0, other]).item()) rng = np.random.default_rng(seed) candidates = [ [index for index, (left, right) in enumerate(slices) if base_mask[left:right, m].any()] for m in range(3) ] masks_to_score: list[np.ndarray] = [] row_specs: list[tuple[float, str, int]] = [] for rate in (0.1, 0.2, 0.3): counts = [max(1, int(np.ceil(rate * len(indices)))) if indices else 0 for indices in candidates] method_choices: list[tuple[str, list[list[int]]]] = [] top_choices: list[list[int]] = [] random_choices: list[list[int]] = [] for modality in range(3): available = candidates[modality] count = min(counts[modality], len(available)) score_order = sorted(available, key=lambda index: (-abs(contribution[modality, index]), index)) top_choices.append(score_order[:count]) random_choices.append(list(rng.choice(available, size=count, replace=False)) if count else []) method_choices.append(("owen", top_choices)) method_choices.append(("matched_random", random_choices)) for method_name, choices in method_choices: reps = 1 if method_name == "owen" else random_replicates for rep in range(reps): if method_name == "matched_random" and rep > 0: choices = [ list(rng.choice(candidates[m], size=min(counts[m], len(candidates[m])), replace=False)) if counts[m] else [] for m in range(3) ] deleted = base_mask.copy() retained = np.zeros_like(base_mask, dtype=bool) for modality, selected in enumerate(choices): for bin_index in selected: left, right = slices[bin_index] deleted[left:right, modality] = False retained[left:right, modality] = base_mask[left:right, modality] masks_to_score.extend((deleted, retained)) row_specs.extend(((rate, method_name, rep), (rate, method_name, rep))) score_values = _margin_values(models, xs, np.stack(masks_to_score), target, other) output: list[dict[str, Any]] = [] cursor = 0 grouped: dict[tuple[float, str], list[tuple[float, float]]] = {} for rate, method_name, _rep in row_specs[::2]: deletion_margin = float(score_values[cursor]) retention_margin = float(score_values[cursor + 1]) cursor += 2 grouped.setdefault((rate, method_name), []).append( (full_margin - deletion_margin, retention_margin) ) for (rate, method_name), values in sorted(grouped.items()): arr = np.asarray(values, dtype=np.float64) output.append( { "sample_id": sample_id, "budget": rate, "method": method_name, "replicates": len(values), "deletion_margin_drop_mean": float(arr[:, 0].mean()), "retention_margin_mean": float(arr[:, 1].mean()), "retention_margin_drop_mean": float((full_margin - arr[:, 1]).mean()), "full_margin": full_margin, } ) return output