Flatten submit package structure

This commit is contained in:
2026-09-26 16:47:17 +08:00
parent 411f0f97e5
commit c116ce60aa
173 changed files with 1242 additions and 1243 deletions
+240
View File
@@ -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