30 lines
1.2 KiB
Python
30 lines
1.2 KiB
Python
from __future__ import annotations
|
|
|
|
import argparse
|
|
import csv
|
|
from pathlib import Path
|
|
|
|
from .train_compare import _plot, _summary, _write_csv
|
|
|
|
|
|
def main() -> None:
|
|
parser = argparse.ArgumentParser(description="Rebuild Q2 summary tables from saved validation predictions")
|
|
parser.add_argument("--output-dir", default=str(Path(__file__).resolve().parents[1] / "outputs" / "followups" / "earlyconcat_standalone"))
|
|
args = parser.parse_args()
|
|
output = Path(args.output_dir)
|
|
with (output / "validation_metrics_by_condition.csv").open(encoding="utf-8-sig", newline="") as stream:
|
|
rows = list(csv.DictReader(stream))
|
|
for row in rows:
|
|
for key in ("missing_rate", "accuracy", "macro_f1", "mae", "pearson", "n_valid"):
|
|
row[key] = float(row[key])
|
|
row["seed"] = int(row["seed"])
|
|
summary = _summary(rows)
|
|
_write_csv(output / "summary.csv", summary)
|
|
aligned = [row for row in summary if row["representation"] == "provided_word_aligned_50"]
|
|
_plot(aligned, rows, output / "missing_rate_comparison.png")
|
|
print(f"rebuilt summary table and plot from {len(rows)} saved validation rows")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|