241 lines
10 KiB
Python
241 lines
10 KiB
Python
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
|