91 lines
4.1 KiB
Python
91 lines
4.1 KiB
Python
"""Plot the aligned-data rate sweep and matched missing-type response."""
|
|
from __future__ import annotations
|
|
|
|
import csv
|
|
from pathlib import Path
|
|
|
|
import matplotlib.pyplot as plt
|
|
import numpy as np
|
|
|
|
|
|
RESULTS = Path(__file__).resolve().parent / "results"
|
|
|
|
|
|
def read_csv(name: str) -> list[dict[str, str]]:
|
|
with (RESULTS / name).open(encoding="utf-8-sig", newline="") as stream:
|
|
return list(csv.DictReader(stream))
|
|
|
|
|
|
def main() -> None:
|
|
rate_rows = read_csv("controlled_missingness.csv")
|
|
rate_bootstrap = read_csv("controlled_group_bootstrap.csv")
|
|
type_bootstrap = read_csv("matched_missing_type_bootstrap.csv")
|
|
|
|
colors = {"single": "#3b82f6", "sync": "#dc2626", "partial": "#16a34a", "async": "#9333ea"}
|
|
labels = {"single": "Single modality", "sync": "Synchronous", "partial": "Partial overlap", "async": "Asynchronous"}
|
|
fig, (ax_rate, ax_type) = plt.subplots(1, 2, figsize=(12.4, 4.8), gridspec_kw={"width_ratios": [1.35, 1.0]})
|
|
|
|
baseline = next(row for row in rate_rows if row["model"] == "C5" and row["mask_pattern"] == "none")
|
|
baseline_ci = next(row for row in rate_bootstrap if row["model"] == "C5" and row["scenario"] == "0.0/none" and row["metric"] == "mae")
|
|
for mode in ("single", "sync", "partial", "async"):
|
|
rows = [baseline] + sorted(
|
|
(row for row in rate_rows if row["model"] == "C5" and row["mask_pattern"] == mode),
|
|
key=lambda row: float(row["rate_requested_per_selected_source"]),
|
|
)
|
|
x, y, lower, upper = [], [], [], []
|
|
for row in rows:
|
|
if row["mask_pattern"] == "none":
|
|
ci = baseline_ci
|
|
scenario = "0.0/none"
|
|
else:
|
|
scenario = f"{float(row['rate_requested_per_selected_source']):.1f}/{mode}"
|
|
ci = next(item for item in rate_bootstrap if item["model"] == "C5" and item["scenario"] == scenario and item["metric"] == "mae")
|
|
x.append(float(row["rate_realized_additional_global"]))
|
|
y.append(float(row["regression_mae"]))
|
|
lower.append(float(ci["ci_2_5"]))
|
|
upper.append(float(ci["ci_97_5"]))
|
|
ax_rate.errorbar(
|
|
x, y, yerr=[np.asarray(y) - np.asarray(lower), np.asarray(upper) - np.asarray(y)],
|
|
color=colors[mode], marker="o", linewidth=1.7, markersize=4.5,
|
|
capsize=2.5, label=labels[mode], alpha=0.95,
|
|
)
|
|
ax_rate.set_title("C5 performance across missing rates")
|
|
ax_rate.set_xlabel("Added missing rate (paper definition)")
|
|
ax_rate.set_ylabel("Regression MAE (95% group-bootstrap CI)")
|
|
ax_rate.grid(axis="both", color="#d1d5db", linewidth=0.7, alpha=0.65)
|
|
ax_rate.legend(frameon=False, fontsize=8.5, loc="upper left")
|
|
|
|
type_order = ("T", "A", "V", "TA", "TV", "AV", "TAV")
|
|
point, low, high = [], [], []
|
|
for label in type_order:
|
|
scenario = f"matched_type_{label}"
|
|
boot = next(row for row in type_bootstrap if row["model"] == "C5" and row["scenario"] == scenario and row["metric"] == "mae")
|
|
point.append(float(boot["delta_to_natural"]))
|
|
low.append(float(boot["delta_to_natural_ci_2_5"]))
|
|
high.append(float(boot["delta_to_natural_ci_97_5"]))
|
|
positions = np.arange(len(type_order))
|
|
ax_type.errorbar(
|
|
positions, point, yerr=[np.asarray(point) - low, high - np.asarray(point)],
|
|
fmt="o", color="#2563eb", ecolor="#2563eb", capsize=3, linewidth=1.4,
|
|
markersize=5,
|
|
)
|
|
ax_type.axhline(0, color="#374151", linewidth=1, linestyle="--")
|
|
ax_type.set_xticks(positions, type_order)
|
|
ax_type.set_title("Matched missing-modality types")
|
|
ax_type.set_xlabel("Hidden modality set")
|
|
ax_type.set_ylabel("MAE change from natural condition")
|
|
ax_type.grid(axis="y", color="#d1d5db", linewidth=0.7, alpha=0.65)
|
|
ax_type.text(
|
|
0.02, 0.02, "Same added feature-row count per sample and type",
|
|
transform=ax_type.transAxes, fontsize=7.5, color="#4b5563",
|
|
)
|
|
|
|
fig.tight_layout(pad=1.2)
|
|
output = RESULTS / "aligned_missingness_effects.png"
|
|
fig.savefig(output, dpi=200, bbox_inches="tight", facecolor="white")
|
|
print(output)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|