143 lines
7.6 KiB
Python
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()
|