Download cas9/evaluate_val_multi_sample.py from ChatterjeeLab/pCoMole: direct link, hf CLI and curl.
- Browser
- Download file 6.51 kB
-
https://huggingface.co/ChatterjeeLab/pCoMole/resolve/main/cas9/evaluate_val_multi_sample.py
- Command line
-
hf download hf://ChatterjeeLab/pCoMole/cas9/evaluate_val_multi_sample.py
-
curl -L -o evaluate_val_multi_sample.py https://huggingface.co/ChatterjeeLab/pCoMole/resolve/main/cas9/evaluate_val_multi_sample.py
6.51 kB
| #!/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() | |