#!/usr/bin/env python3 from __future__ import annotations import argparse import json from pathlib import Path import pandas as pd from bgc_retrieval.artifacts import sha256_file, write_json_immutable from bgc_retrieval.statistics import holm_adjust, paired_family_test METRICS = ["recall@50", "mrr", "map", "ndcg@50"] def normalized_method(method: str) -> str: if method.startswith("residual_pfam_validation_"): return "residual_pfam_validation_selected" if method.startswith("residual_validation_alpha_"): return "residual_validation_selected" return method def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("--artifact-root", default="artifacts") parser.add_argument("--seed", action="append", type=int, required=True) parser.add_argument("--output-dir", required=True) parser.add_argument("--ensemble-run") args = parser.parse_args() root = Path(args.artifact_root) output = Path(args.output_dir) output.mkdir(parents=True, exist_ok=False) seed_rows = [] paired_rows = [] inputs: list[Path] = [] for seed in args.seed: training_dir = root / f"paper-v2-residual-seed-{seed}" evaluation_dir = root / f"paper-v2-residual-seed-{seed}-evaluation" history_path = training_dir / "residual_history.json" metadata_path = evaluation_dir / "metadata.json" summary_path = evaluation_dir / "summary.csv" groups_path = evaluation_dir / "group_results.csv" inputs.extend([history_path, metadata_path, summary_path, groups_path]) history = json.loads(history_path.read_text(encoding="utf-8")) best = max(history, key=lambda row: row["validation_recall@50"]) metadata = json.loads(metadata_path.read_text(encoding="utf-8")) summary = pd.read_csv(summary_path) for record in summary.to_dict("records"): seed_rows.append( { "training_seed": seed, "method": normalized_method(str(record["method"])), "metric": str(record["metric"]), "value": float(record["mean"]), "selected_residual_alpha": metadata["selected_residual_alpha"], "selected_hybrid_alpha": metadata["selected_hybrid_alpha"], "selected_pfam_beta": metadata["selected_pfam_beta"], "best_epoch": int(best["epoch"]), "best_validation_recall@50": float(best["validation_recall@50"]), "epochs": len(history), } ) groups = pd.read_csv(groups_path) actual = { normalized_method(str(method)): str(method) for method in groups["method"].unique() } comparisons = [ ("weighted_vs_raw", actual["weighted_gene_esm"], "raw_esm_mean"), ( "residual_vs_raw", actual["residual_validation_selected"], "raw_esm_mean", ), ( "hybrid_vs_pfam", actual["residual_pfam_validation_selected"], "pfam_jaccard_max", ), ] for family, method, baseline in comparisons: results = [ paired_family_test(groups, method, baseline, metric) for metric in METRICS ] adjusted = holm_adjust(row["p_value"] for row in results) for row, corrected in zip(results, adjusted): paired_rows.append( { "training_seed": seed, "comparison": family, **row, "p_value_holm": corrected, } ) seed_frame = pd.DataFrame(seed_rows).sort_values( ["method", "metric", "training_seed"] ) seed_frame.to_csv(output / "seed_level_summary.csv", index=False) paired_frame = pd.DataFrame(paired_rows).sort_values( ["comparison", "metric", "training_seed"] ) paired_frame.to_csv(output / "paired_comparisons.csv", index=False) aggregate = ( seed_frame.groupby(["method", "metric"], as_index=False)["value"] .agg( n="count", mean="mean", sample_std="std", minimum="min", maximum="max", ) .sort_values(["method", "metric"]) ) aggregate.to_csv(output / "aggregate_summary.csv", index=False) ensemble_record = None if args.ensemble_run: ensemble_dir = root / args.ensemble_run ensemble_summary = ensemble_dir / "summary.csv" ensemble_metadata = ensemble_dir / "metadata.json" inputs.extend([ensemble_summary, ensemble_metadata]) pd.read_csv(ensemble_summary).to_csv( output / "ensemble_summary.csv", index=False ) ensemble_record = json.loads( ensemble_metadata.read_text(encoding="utf-8") ) write_json_immutable( output / "analysis_metadata.json", { "schema_version": 1, "organism_scope": "Streptomyces griseus", "analysis_status": "post_hoc_redesign_pilot", "seeds": args.seed, "ensemble_metadata": ensemble_record, "input_sha256": { str(path.relative_to(root)): sha256_file(path) for path in sorted(inputs) }, }, ) print(output) if __name__ == "__main__": main()