| |
| 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() |
|
|