#!/usr/bin/env python3 """ Run evaluate_val-style evaluation multiple times and report per-run metric lists. Stochastic steps (sampling from the test set, flow time / noise, generation) can make scalar metrics differ between runs. This script calls run_evaluate_val() repeatedly and prints every value plus simple mean / std / min / max summaries. Seeding: --base_seed unset: same as legacy evaluate_val (no explicit RNG seeding; variation across runs). --base_seed K: run i uses seed K + i before each replicate (reproducible multi-run bracket). """ import argparse import json import numpy as np from cas9.evaluate_val import get_evaluate_val_argument_parser, run_evaluate_val, print_evaluation_results def _get_nested(d, *keys): for k in keys: d = d[k] return d def _stats(vals): arr = np.array([v for v in vals if v is not None], dtype=float) if arr.size == 0: return None return { "mean": float(arr.mean()), "std": float(arr.std()), "min": float(arr.min()), "max": float(arr.max()), } def _fmt_list(vals, prec=4): parts = [] for v in vals: if v is None: parts.append("None") elif isinstance(v, float): parts.append(f"{v:.{prec}f}") elif isinstance(v, int): parts.append(str(v)) else: parts.append(str(v)) return "[" + ", ".join(parts) + "]" def _print_scalar_block(title, rows, all_metrics): print(f"\n{title}") print("-" * len(title)) for label, path in rows: vals = [_get_nested(m, *path) for m in all_metrics] print(f" {label}") print(f" values: {_fmt_list(vals)}") st = _stats(vals) if st: print(f" mean={st['mean']:.6f} std={st['std']:.6f} min={st['min']:.6f} max={st['max']:.6f}") else: print(" (no numeric values)") def print_multi_run_summary(all_metrics, args): n = len(all_metrics) print("\n" + "=" * 70) print(f"MULTI-RUN SUMMARY ({n} runs)") print("=" * 70) seeds = [m["run_seed"] for m in all_metrics] print(f"Per-run seeds: {seeds}") loss_rows = [ ("Train avg (full train; None if not computed)", ("loss", "train_avg")), ("Decoded avg (full test/val)", ("loss", "decoded_full_avg")), ("Decoded avg (sampled)", ("loss", "decoded_sampled_avg")), ("Generated avg", ("loss", "generated_avg")), ] _print_scalar_block("SEQUENCE LOSS (val_unweighted_total_loss)", loss_rows, all_metrics) if all_metrics[0]["cas9"] is not None: cas9_rows = [ ("Decoded validity rate", ("cas9", "decoded_validity_rate")), ("Decoded avg Cas9 score", ("cas9", "decoded_avg_score")), ("Generated validity rate", ("cas9", "generated_validity_rate")), ("Generated avg Cas9 score", ("cas9", "generated_avg_score")), ] _print_scalar_block("CAS9 SCORES", cas9_rows, all_metrics) div_dec = [ ("unique_count", ("diversity_decoded", "unique_count")), ("uniqueness_ratio", ("diversity_decoded", "uniqueness_ratio")), ("kmer_diversity", ("diversity_decoded", "kmer_diversity")), ("kmer_avg_similarity", ("diversity_decoded", "kmer_avg_similarity")), ("levenshtein_diversity", ("diversity_decoded", "levenshtein_diversity")), ("levenshtein_avg_similarity", ("diversity_decoded", "levenshtein_avg_similarity")), ] _print_scalar_block("DIVERSITY — decoded (sampled)", div_dec, all_metrics) div_gen = [ ("unique_count", ("diversity_generated", "unique_count")), ("uniqueness_ratio", ("diversity_generated", "uniqueness_ratio")), ("kmer_diversity", ("diversity_generated", "kmer_diversity")), ("kmer_avg_similarity", ("diversity_generated", "kmer_avg_similarity")), ("levenshtein_diversity", ("diversity_generated", "levenshtein_diversity")), ("levenshtein_avg_similarity", ("diversity_generated", "levenshtein_avg_similarity")), ] _print_scalar_block("DIVERSITY — generated", div_gen, all_metrics) if any(m.get("plddt") is not None for m in all_metrics): p_rows = [ ("Decoded mean pLDDT", ("plddt", "decoded_mean")), ("Decoded std pLDDT", ("plddt", "decoded_std")), ("Generated mean pLDDT", ("plddt", "generated_mean")), ("Generated std pLDDT", ("plddt", "generated_std")), ] _print_scalar_block("PLDDT", p_rows, all_metrics) print("=" * 70) def main(): parser = get_evaluate_val_argument_parser() parser.add_argument( "--n_runs", type=int, default=3, help="Number of independent evaluate_val runs (default: 3)", ) parser.add_argument( "--base_seed", type=int, default=None, help="If set, run i uses RNG seed (base_seed + i). If unset, do not seed (like default evaluate_val).", ) parser.add_argument( "--save_fasta", action="store_true", help="Write generated_sequences_run{i}.fasta per run under output_dir", ) parser.add_argument( "--print_each_run", action="store_true", help="Print the full EVALUATION RESULTS block after every replicate", ) parser.add_argument( "--json_out", type=str, default=None, help="Optional path to write a JSON list of per-run metric dicts", ) args = parser.parse_args() if args.n_runs < 1: raise ValueError("--n_runs must be >= 1") all_metrics = [] for i in range(args.n_runs): print("\n" + "#" * 70) print(f"RUN {i + 1} / {args.n_runs}") print("#" * 70) run_seed = (args.base_seed + i) if args.base_seed is not None else None if run_seed is not None: print(f"(run_seed={run_seed})") fasta_tag = f"run{i + 1}" if args.save_fasta else None metrics = run_evaluate_val( args, run_seed=run_seed, save_fasta=args.save_fasta, fasta_tag=fasta_tag, ) all_metrics.append(metrics) if args.print_each_run: print_evaluation_results(metrics, args) print_multi_run_summary(all_metrics, args) if args.json_out: with open(args.json_out, "w") as f: json.dump(all_metrics, f, indent=2) print(f"\nWrote per-run metrics to {args.json_out}") if __name__ == "__main__": main()