"""Create the method-comparison figures from Q2's fixed evaluation outputs.""" from __future__ import annotations import csv import json from pathlib import Path import matplotlib.pyplot as plt import numpy as np from train import RESULTS def read_csv(path: Path) -> list[dict[str, str]]: with path.open("r", encoding="utf-8-sig", newline="") as stream: return list(csv.DictReader(stream)) def main() -> None: controlled = read_csv(RESULTS / "controlled_missingness.csv") test = read_csv(RESULTS / "test_predictions.csv") gates = read_csv(RESULTS / "test_gate_diagnostics.csv") metrics = json.loads((RESULTS / "test_metrics.json").read_text(encoding="utf-8")) selected = str(metrics["selected_model"]) figure, axes = plt.subplots(2, 3, figsize=(17, 10), constrained_layout=True) rate_axis, modality_axis, location_axis, confusion_axis, scatter_axis, interval_axis = axes.flat comparison_models = ("C0", "C3", "C4", "C5", "C6", "C7_distill", "C7_group") for model in comparison_models: subset = [row for row in controlled if row["model"] == model and (row["mask_pattern"] == "none" or row["mask_pattern"] == "single")] subset.sort(key=lambda row: float(row["rate_realized_additional_global"])) if subset: rate_axis.plot([float(row["rate_realized_additional_global"]) for row in subset], [float(row["regression_mae"]) for row in subset], marker="o", label=model) rate_axis.set(title="MAE by realized additional missing rate", xlabel="Additional missing rate (equal T/A/V)", ylabel="MAE") rate_axis.legend(fontsize=8, ncol=2) rate_axis.grid(alpha=0.25) modality_labels = ("T", "A", "V", "TA", "TV", "AV", "TAV") modality_scenarios = {f"0.3/modality_{label}": label for label in modality_labels} modality_models = ("C0", "C5", "C6", "C7_distill", "C7_group") modality_values = np.full((len(modality_models), len(modality_labels)), np.nan) for i, model in enumerate(modality_models): for j, (scenario, _) in enumerate(modality_scenarios.items()): pattern = scenario.split("/", 1)[1] row = next((r for r in controlled if r["model"] == model and r["mask_pattern"] == pattern), None) if row is not None: modality_values[i, j] = float(row["regression_mae"]) image = modality_axis.imshow(modality_values, aspect="auto", cmap="viridis") modality_axis.set(title="Modality combination control: MAE", xticks=range(len(modality_labels)), xticklabels=modality_labels, yticks=range(len(modality_models)), yticklabels=modality_models) modality_axis.tick_params(axis="x", rotation=35) figure.colorbar(image, ax=modality_axis, fraction=0.046, pad=0.04) position_values = np.full((3, 3), np.nan) for i, modality in enumerate(("T", "A", "V")): for j, location in enumerate(("start", "middle", "end")): scenario = f"0.3/location_{location}_{modality}" pattern = scenario.split("/", 1)[1] row = next((r for r in controlled if r["model"] == selected and r["mask_pattern"] == pattern), None) if row is not None: position_values[i, j] = float(row["regression_mae"]) image = location_axis.imshow(position_values, aspect="auto", cmap="magma") location_axis.set(title=f"Selected model {selected}: location MAE", xticks=range(3), xticklabels=("start", "middle", "end"), yticks=range(3), yticklabels=("T", "A", "V")) figure.colorbar(image, ax=location_axis, fraction=0.046, pad=0.04) confusion = np.zeros((3, 3), dtype=np.int64) for row in test: confusion[int(row["true_class"]), int(row["predicted_class"])] += 1 image = confusion_axis.imshow(confusion, cmap="Blues") for i in range(3): for j in range(3): confusion_axis.text(j, i, str(confusion[i, j]), ha="center", va="center") confusion_axis.set(title=f"Test confusion matrix: {selected}", xlabel="Predicted", ylabel="True", xticks=range(3), xticklabels=("negative", "neutral", "positive"), yticks=range(3), yticklabels=("negative", "neutral", "positive")) figure.colorbar(image, ax=confusion_axis, fraction=0.046, pad=0.04) true_score = np.asarray([float(row["true_sentiment"]) for row in test]) predicted_score = np.asarray([float(row["predicted_sentiment"]) for row in test]) scatter_axis.scatter(true_score, predicted_score, alpha=0.55, s=18) scatter_axis.plot([-3, 3], [-3, 3], "k--", linewidth=1) scatter_axis.set(title=f"Test sentiment: MAE={metrics['regression_mae']:.3f}", xlabel="True sentiment", ylabel="Predicted sentiment", xlim=(-3, 3), ylim=(-3, 3)) scatter_axis.grid(alpha=0.2) lower = np.asarray([float(row["interval_90_lower"]) for row in test]) upper = np.asarray([float(row["interval_90_upper"]) for row in test]) width = upper - lower covered = (true_score >= lower) & (true_score <= upper) order = np.argsort(width) bins = np.array_split(order, min(10, len(order))) interval_axis.plot([width[idx].mean() for idx in bins], [covered[idx].mean() for idx in bins], marker="o") interval_axis.axhline(0.9, color="black", linestyle="--", linewidth=1, label="nominal 90%") interval_axis.set(title="Test interval coverage by width decile", xlabel="Mean interval width", ylabel="Empirical coverage", ylim=(0, 1)) interval_axis.legend() interval_axis.grid(alpha=0.2) figure.suptitle("Q2 validation controls and official-test diagnostics", fontsize=15) figure.savefig(RESULTS / "q2_diagnostics.png", dpi=160) plt.close(figure) steps = 50 modalities = ("text", "audio", "vision") weight_sum = np.zeros((steps, len(modalities)), dtype=np.float64) reliability_sum = np.zeros_like(weight_sum) count = np.zeros_like(weight_sum) for row in gates: if row["modality"] not in modalities: continue t, m = int(row["step"]), modalities.index(row["modality"]) weight_sum[t, m] += float(row["fusion_weight_mean_over_paths"]) reliability_sum[t, m] += float(row["reliability"]) count[t, m] += 1 weights = weight_sum / np.maximum(count, 1.0) reliabilities = reliability_sum / np.maximum(count, 1.0) gate_figure, gate_axis = plt.subplots(figsize=(12, 5), constrained_layout=True) for m, modality in enumerate(modalities): gate_axis.plot(range(steps), weights[:, m], label=f"{modality} fusion weight") gate_axis.set(title=f"Test mean fusion gates by position: {selected}", xlabel="Aligned step", ylabel="Mean fusion weight") gate_axis.legend(ncol=3) gate_axis.grid(alpha=0.25) reliability_axis = gate_axis.twinx() for m, modality in enumerate(modalities): reliability_axis.plot(range(steps), reliabilities[:, m], linestyle=":", alpha=0.7, label=f"{modality} reliability") reliability_axis.set_ylabel("Mean reliability proxy") handles, labels = gate_axis.get_legend_handles_labels() right_handles, right_labels = reliability_axis.get_legend_handles_labels() gate_axis.legend(handles + right_handles, labels + right_labels, ncol=3, fontsize=8) gate_figure.savefig(RESULTS / "q2_gate_positions.png", dpi=160) plt.close(gate_figure) manifest_path = RESULTS / "run_manifest.json" manifest = json.loads(manifest_path.read_text(encoding="utf-8")) manifest["diagnostic_figures"] = ["q2_diagnostics.png", "q2_gate_positions.png"] manifest_path.write_text(json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8") print(f"Wrote Q2 figures for selected model {selected} to {RESULTS}") if __name__ == "__main__": main()