Prepare minimum submission bundle
This commit is contained in:
@@ -0,0 +1,240 @@
|
||||
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
|
||||
Reference in New Issue
Block a user