pCoMole / cas9 /evaluate_val_multi_sample.py
Maximilian Holsman
Claude Opus 5
Add Cas9 task
12fea4a
Raw History Blame Contribute Delete
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()