Files
modeling_zhaocui/deep_learning/Q2/q2/analyze_aligned_missingness.py

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)