File size: 6,508 Bytes
12fea4a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
#!/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()