362 lines
17 KiB
Python
362 lines
17 KiB
Python
"""Factorial local-missingness analysis for the aligned Q2 checkpoints."""
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import csv
|
|
import json
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
import matplotlib.pyplot as plt
|
|
import numpy as np
|
|
import torch
|
|
|
|
from .data import ATTACHMENT2, RobustStats, apply_robust_stats, load_aligned
|
|
from .evaluate_math_protocol import (
|
|
SCENARIO_SEED,
|
|
continuous_mask,
|
|
metrics,
|
|
scenario_seed,
|
|
sha256,
|
|
write_csv,
|
|
)
|
|
from .models import AlignedFusionModel
|
|
from .mofe import MixtureOfFusionExperts
|
|
from .train_mofe import EARLYCONCAT, MODEL_CONFIG, MOFE7_MLP, _predict
|
|
|
|
|
|
Q2_ROOT = Path(__file__).resolve().parents[1]
|
|
RUN_DIR = Q2_ROOT / "outputs" / "followups" / "R03_math_protocol_retraining"
|
|
OUTPUT_DIR = RUN_DIR / "aligned_missingness_analysis"
|
|
SEED = 20260924
|
|
BOOTSTRAP_SEED = 20260927
|
|
MISSING_RATES = (0.1, 0.3, 0.5, 0.7)
|
|
MODALITY_SETS: dict[str, tuple[int, ...]] = {
|
|
"T": (0,),
|
|
"A": (1,),
|
|
"V": (2,),
|
|
"TA": (0, 1),
|
|
"TV": (0, 2),
|
|
"AV": (1, 2),
|
|
"TAV": (0, 1, 2),
|
|
}
|
|
METHODS = (EARLYCONCAT, MOFE7_MLP)
|
|
METRICS = ("accuracy", "macro_f1", "mae", "pearson")
|
|
|
|
|
|
def _device(name: str) -> torch.device:
|
|
if name == "auto":
|
|
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
return torch.device(name)
|
|
|
|
|
|
def _load_model(method: str, dims: tuple[int, int, int], device: torch.device) -> torch.nn.Module:
|
|
checkpoint_path = RUN_DIR / "models" / method / f"seed_{SEED}" / "model_best.pt"
|
|
state = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
|
if method == EARLYCONCAT:
|
|
if state.get("kind") != "concat":
|
|
raise ValueError(f"unexpected EarlyConcat checkpoint format: {checkpoint_path}")
|
|
model: torch.nn.Module = AlignedFusionModel("concat", dims=dims).to(device)
|
|
elif method == MOFE7_MLP:
|
|
if state.get("config") != MODEL_CONFIG:
|
|
raise ValueError(f"unexpected MoFE checkpoint configuration: {checkpoint_path}")
|
|
model = MixtureOfFusionExperts(dims=dims, **MODEL_CONFIG).to(device)
|
|
else:
|
|
raise ValueError(f"unknown model: {method}")
|
|
if int(state.get("seed", -1)) != SEED or tuple(state.get("dims", ())) != dims:
|
|
raise ValueError(f"checkpoint metadata mismatch: {checkpoint_path}")
|
|
model.load_state_dict(state["state_dict"])
|
|
return model.eval()
|
|
|
|
|
|
def _scenario_key(label: str, rate: float) -> str:
|
|
return f"{label}/{rate:.1f}/middle_sync"
|
|
|
|
|
|
def _factorial_masks(valid) -> tuple[dict[str, np.ndarray], dict[str, dict[str, float]]]:
|
|
scenarios = {"clean": valid.mask.copy()}
|
|
realized: dict[str, dict[str, float]] = {}
|
|
for label, selected in MODALITY_SETS.items():
|
|
for rate in MISSING_RATES:
|
|
key = _scenario_key(label, rate)
|
|
mask_rows = []
|
|
for sample_id, original in zip(valid.ids, valid.mask):
|
|
rng = np.random.default_rng(scenario_seed(SCENARIO_SEED, sample_id, key))
|
|
corrupted = continuous_mask(
|
|
original,
|
|
rate,
|
|
"single",
|
|
rng,
|
|
modalities=selected,
|
|
location="middle",
|
|
)
|
|
mask_rows.append(corrupted)
|
|
current = np.stack(mask_rows)
|
|
scenarios[key] = current
|
|
# Mean realized missing fraction among selected modalities. The
|
|
# unselected modalities are deliberately excluded from this rate.
|
|
per_sample = []
|
|
for original, corrupted in zip(valid.mask, current):
|
|
before = original[:, selected].sum(axis=0)
|
|
hidden = (original[:, selected] & ~corrupted[:, selected]).sum(axis=0)
|
|
rates = np.divide(hidden, before, out=np.full(len(selected), np.nan), where=before > 0)
|
|
if np.isfinite(rates).any():
|
|
per_sample.append(float(np.nanmean(rates)))
|
|
realized[key] = {
|
|
"requested_rate": float(rate),
|
|
"selected_modality_rate_mean": float(np.mean(per_sample)) if per_sample else float("nan"),
|
|
"selected_modality_rate_min": float(np.min(per_sample)) if per_sample else float("nan"),
|
|
"selected_modality_rate_max": float(np.max(per_sample)) if per_sample else float("nan"),
|
|
}
|
|
return scenarios, realized
|
|
|
|
|
|
def _condition_summary_rows(
|
|
valid,
|
|
predictions: dict[tuple[str, str], dict[str, np.ndarray]],
|
|
realized: dict[str, dict[str, float]],
|
|
) -> list[dict[str, Any]]:
|
|
rows: list[dict[str, Any]] = []
|
|
clean_key = "clean"
|
|
for label in MODALITY_SETS:
|
|
for rate in MISSING_RATES:
|
|
scenario = _scenario_key(label, rate)
|
|
for method in METHODS:
|
|
pred = predictions[(method, scenario)]
|
|
clean = predictions[(method, clean_key)]
|
|
current_metrics = metrics(valid, pred["logits"], pred["intensity"])
|
|
clean_metrics = metrics(valid, clean["logits"], clean["intensity"])
|
|
rows.append({
|
|
"method": method,
|
|
"missing_modalities": label,
|
|
"selected_modalities": "+".join(label),
|
|
"requested_missing_rate": rate,
|
|
"missing_layout": "centered contiguous span per selected modality; other modalities unchanged",
|
|
"realized_selected_modality_rate_mean": realized[scenario]["selected_modality_rate_mean"],
|
|
"realized_selected_modality_rate_min": realized[scenario]["selected_modality_rate_min"],
|
|
"realized_selected_modality_rate_max": realized[scenario]["selected_modality_rate_max"],
|
|
"n_valid": valid.n,
|
|
**current_metrics,
|
|
"delta_accuracy_vs_clean": current_metrics["accuracy"] - clean_metrics["accuracy"],
|
|
"delta_macro_f1_vs_clean": current_metrics["macro_f1"] - clean_metrics["macro_f1"],
|
|
"delta_mae_vs_clean": current_metrics["mae"] - clean_metrics["mae"],
|
|
"delta_pearson_vs_clean": current_metrics["pearson"] - clean_metrics["pearson"],
|
|
})
|
|
return rows
|
|
|
|
|
|
def _aggregate_rows(rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
|
result: list[dict[str, Any]] = []
|
|
methods = list(METHODS)
|
|
for method in methods:
|
|
for rate in MISSING_RATES:
|
|
subset = [r for r in rows if r["method"] == method and r["requested_missing_rate"] == rate]
|
|
result.append({
|
|
"method": method,
|
|
"requested_missing_rate": rate,
|
|
"n_modality_sets": len(subset),
|
|
**{metric: float(np.mean([float(r[metric]) for r in subset])) for metric in METRICS},
|
|
"mean_macro_f1_drop_vs_clean": float(-np.mean([float(r["delta_macro_f1_vs_clean"]) for r in subset])),
|
|
"mean_mae_increase_vs_clean": float(np.mean([float(r["delta_mae_vs_clean"]) for r in subset])),
|
|
})
|
|
return result
|
|
|
|
|
|
def _ablation_rows(rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
|
result: list[dict[str, Any]] = []
|
|
for row in rows:
|
|
other_method = MOFE7_MLP if row["method"] == EARLYCONCAT else EARLYCONCAT
|
|
other = next(r for r in rows if r["method"] == other_method
|
|
and r["missing_modalities"] == row["missing_modalities"]
|
|
and r["requested_missing_rate"] == row["requested_missing_rate"])
|
|
if row["method"] != EARLYCONCAT:
|
|
continue
|
|
result.append({
|
|
"missing_modalities": row["missing_modalities"],
|
|
"requested_missing_rate": row["requested_missing_rate"],
|
|
"delta_macro_f1_mofe_minus_earlyconcat": float(other["macro_f1"] - row["macro_f1"]),
|
|
"delta_accuracy_mofe_minus_earlyconcat": float(other["accuracy"] - row["accuracy"]),
|
|
"delta_mae_mofe_minus_earlyconcat": float(other["mae"] - row["mae"]),
|
|
"delta_pearson_mofe_minus_earlyconcat": float(other["pearson"] - row["pearson"]),
|
|
})
|
|
return result
|
|
|
|
|
|
def _plot_rate_curves(rows: list[dict[str, Any]], clean_metrics: dict[str, dict[str, float]], output_dir: Path) -> None:
|
|
colors = {"T": "#3366cc", "A": "#dc3912", "V": "#ff9900", "TA": "#109618", "TV": "#990099", "AV": "#0099c6", "TAV": "#dd4477"}
|
|
fig, axes = plt.subplots(1, 2, figsize=(13, 5), sharey=True)
|
|
for ax, method in zip(axes, METHODS):
|
|
for label in MODALITY_SETS:
|
|
subset = sorted(
|
|
(r for r in rows if r["method"] == method and r["missing_modalities"] == label),
|
|
key=lambda r: r["requested_missing_rate"],
|
|
)
|
|
xs = [float(r["requested_missing_rate"]) for r in subset]
|
|
ys = [float(r["macro_f1"]) for r in subset]
|
|
ax.plot(xs, ys, marker="o", linewidth=1.8, label=label, color=colors[label])
|
|
base = clean_metrics[method]["macro_f1"]
|
|
ax.axhline(base, color="#333333", linestyle="--", linewidth=1.2, label="clean")
|
|
ax.set_title(method)
|
|
ax.set_xlabel("Requested missing rate of selected modalities")
|
|
ax.set_xticks(MISSING_RATES)
|
|
ax.grid(alpha=0.25)
|
|
axes[0].set_ylabel("Validation Macro-F1")
|
|
axes[1].legend(title="Missing set", bbox_to_anchor=(1.02, 1), loc="upper left")
|
|
fig.suptitle("Aligned Q2: local missingness rate and modality type")
|
|
fig.tight_layout()
|
|
fig.savefig(output_dir / "macro_f1_by_modality_and_rate.png", dpi=180, bbox_inches="tight")
|
|
plt.close(fig)
|
|
|
|
fig, axes = plt.subplots(1, 2, figsize=(12, 5), sharey=True)
|
|
for ax, method in zip(axes, METHODS):
|
|
matrix = np.asarray([
|
|
[next(float(r["delta_macro_f1_vs_clean"]) for r in rows
|
|
if r["method"] == method and r["missing_modalities"] == label
|
|
and r["requested_missing_rate"] == rate)
|
|
for rate in MISSING_RATES]
|
|
for label in MODALITY_SETS
|
|
])
|
|
image = ax.imshow(matrix, aspect="auto", cmap="RdYlGn", vmin=-0.12, vmax=0.04)
|
|
ax.set_title(method)
|
|
ax.set_xticks(range(len(MISSING_RATES)), [f"{int(r*100)}%" for r in MISSING_RATES])
|
|
ax.set_yticks(range(len(MODALITY_SETS)), list(MODALITY_SETS))
|
|
ax.set_xlabel("Requested missing rate")
|
|
for i in range(matrix.shape[0]):
|
|
for j in range(matrix.shape[1]):
|
|
ax.text(j, i, f"{matrix[i, j]:+.3f}", ha="center", va="center", fontsize=8)
|
|
axes[0].set_ylabel("Selected modality set")
|
|
fig.colorbar(image, ax=axes.ravel().tolist(), label="Macro-F1 change vs clean")
|
|
fig.suptitle("Aligned Q2: Macro-F1 change under modality ablation")
|
|
fig.savefig(output_dir / "macro_f1_drop_heatmap.png", dpi=180, bbox_inches="tight")
|
|
plt.close(fig)
|
|
|
|
|
|
def _plot_location_span(location_rows: list[dict[str, Any]], output_dir: Path) -> None:
|
|
locations = ("start", "middle", "end")
|
|
fig, axes = plt.subplots(1, 2, figsize=(12, 4.5), sharey=True)
|
|
for ax, method in zip(axes, METHODS):
|
|
for modality in ("T", "A", "V"):
|
|
values = []
|
|
for location in locations:
|
|
row = next(r for r in location_rows if r["method"] == method
|
|
and r["kind"] == "location" and r["label"] == modality
|
|
and r["variant"] == location)
|
|
values.append(float(row["macro_f1"]))
|
|
ax.plot(locations, values, marker="o", label=modality)
|
|
ax.set_title(method)
|
|
ax.set_ylabel("Validation Macro-F1")
|
|
ax.set_xlabel("30% missing-block location")
|
|
ax.grid(alpha=0.25)
|
|
axes[1].legend(title="Modality")
|
|
fig.suptitle("Aligned Q2: sensitivity to missing-block location")
|
|
fig.tight_layout()
|
|
fig.savefig(output_dir / "location_sensitivity.png", dpi=180, bbox_inches="tight")
|
|
plt.close(fig)
|
|
|
|
|
|
def run(device_name: str = "auto", output_dir: Path = OUTPUT_DIR) -> None:
|
|
output_dir.mkdir(parents=True, exist_ok=True)
|
|
device = _device(device_name)
|
|
if device.type == "cuda" and not torch.cuda.is_available():
|
|
raise RuntimeError("CUDA requested but not available")
|
|
feature_path = ATTACHMENT2 / "aligned_50.pkl"
|
|
splits = load_aligned(feature_path)
|
|
valid = apply_robust_stats(splits["valid"], RobustStats.load(RUN_DIR / "aligned_robust_stats.npz"))
|
|
dims = tuple(int(x.shape[-1]) for x in valid.x)
|
|
scenarios, realized = _factorial_masks(valid)
|
|
predictions: dict[tuple[str, str], dict[str, np.ndarray]] = {}
|
|
clean_metrics: dict[str, dict[str, float]] = {}
|
|
for method in METHODS:
|
|
model = _load_model(method, dims, device)
|
|
for scenario, mask in scenarios.items():
|
|
predictions[(method, scenario)] = _predict(model, valid, mask, device, batch_size=64)
|
|
clean_metrics[method] = metrics(
|
|
valid,
|
|
predictions[(method, "clean")]["logits"],
|
|
predictions[(method, "clean")]["intensity"],
|
|
)
|
|
del model
|
|
if torch.cuda.is_available():
|
|
torch.cuda.empty_cache()
|
|
|
|
rows = _condition_summary_rows(valid, predictions, realized)
|
|
aggregate_rows = _aggregate_rows(rows)
|
|
ablation_rows = _ablation_rows(rows)
|
|
write_csv(output_dir / "modality_rate_metrics.csv", rows)
|
|
write_csv(output_dir / "modality_rate_summary.csv", aggregate_rows)
|
|
write_csv(output_dir / "architecture_ablation_deltas.csv", ablation_rows)
|
|
|
|
# Analyze the existing fixed 30% location, span, and synchrony controls.
|
|
old_conditions_path = RUN_DIR / "controlled_metrics_by_scenario.csv"
|
|
with old_conditions_path.open("r", newline="", encoding="utf-8-sig") as stream:
|
|
old_rows = list(csv.DictReader(stream))
|
|
location_rows: list[dict[str, Any]] = []
|
|
for row in old_rows:
|
|
scenario = row["scenario"]
|
|
pieces = scenario.split("/")
|
|
if len(pieces) != 2 or pieces[0] != "0.3":
|
|
continue
|
|
subparts = pieces[1].split("_")
|
|
if subparts[0] == "location":
|
|
kind, variant, label = "location", subparts[1], subparts[2]
|
|
elif subparts[0] == "span":
|
|
kind, variant, label = "span", subparts[1], subparts[2]
|
|
elif subparts[0] == "synchrony":
|
|
kind, variant, label = "synchrony", subparts[1], "TAV"
|
|
else:
|
|
continue
|
|
location_rows.append({
|
|
"method": row["method"],
|
|
"kind": kind,
|
|
"variant": variant,
|
|
"label": label,
|
|
"macro_f1": float(row["macro_f1"]),
|
|
"accuracy": float(row["accuracy"]),
|
|
"mae": float(row["mae"]),
|
|
"pearson": float(row["pearson"]),
|
|
"scenario": scenario,
|
|
})
|
|
write_csv(output_dir / "location_span_synchrony_metrics.csv", location_rows)
|
|
_plot_rate_curves(rows, clean_metrics, output_dir)
|
|
_plot_location_span(location_rows, output_dir)
|
|
|
|
manifest = {
|
|
"experiment": "Factorial local missingness type x rate on supplied aligned_50 validation data",
|
|
"feature_file": str(feature_path),
|
|
"feature_sha256": sha256(feature_path),
|
|
"representation": "aligned_50 ordered wordpiece positions, not physical-time bins",
|
|
"split": "official validation only",
|
|
"n_valid": valid.n,
|
|
"source_video_groups": len({sid.split("$_$", 1)[0] for sid in valid.ids}),
|
|
"checkpoint_source": str(RUN_DIR / "models"),
|
|
"seed": SEED,
|
|
"device": str(device),
|
|
"cuda_device": torch.cuda.get_device_name(0) if device.type == "cuda" else None,
|
|
"missing_rate_grid": list(MISSING_RATES),
|
|
"missing_modality_sets": {k: list(v) for k, v in MODALITY_SETS.items()},
|
|
"factorial_design": "28 local-missingness conditions plus clean reference; each selected modality receives a centered contiguous span; other modalities remain unchanged",
|
|
"missing_rate_definition": "newly hidden observed positions divided by originally observed positions, averaged over selected modalities; realized rate reported per condition",
|
|
"preserve_at_least_fraction": 0.2,
|
|
"scenario_seed": SCENARIO_SEED,
|
|
"additional_existing_controls": "30% start/middle/end location, one-long/multiple-short span, and sync/partial/async controls imported from the R03 42-scenario audit",
|
|
"test_split_read_or_evaluated": False,
|
|
"label_usage": "validation labels used only for metric computation; no model fitting or checkpoint selection in this analysis",
|
|
}
|
|
(output_dir / "run_manifest.json").write_text(json.dumps(manifest, indent=2), encoding="utf-8")
|
|
print(f"saved aligned missingness analysis to {output_dir}", flush=True)
|
|
print(f"conditions={len(rows)}; methods={len(METHODS)}; device={device}", flush=True)
|
|
for row in aggregate_rows:
|
|
print(
|
|
f"{row['method']} rate={row['requested_missing_rate']:.1f} "
|
|
f"F1={row['macro_f1']:.4f} MAE={row['mae']:.4f}",
|
|
flush=True,
|
|
)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
parser = argparse.ArgumentParser(description=__doc__)
|
|
parser.add_argument("--device", default="auto", choices=("auto", "cuda", "cpu"))
|
|
parser.add_argument("--output-dir", type=Path, default=OUTPUT_DIR)
|
|
args = parser.parse_args()
|
|
run(device_name=args.device, output_dir=args.output_dir)
|