整理 Q1-Q3 实验代码与结果
This commit is contained in:
@@ -0,0 +1,29 @@
|
||||
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" / "algorithm_selection"))
|
||||
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()
|
||||
Reference in New Issue
Block a user