Files

143 lines
7.6 KiB
Python

"""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()