File size: 7,507 Bytes
7d3c9bf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""
generate_valid.py -- standalone CLI to generate peptide SMILES with a chosen
validity-boosting sampling strategy (see ``sampling_strategies.py``) and report the
fraction that pass ``utils.app.PeptideAnalyzer.is_peptide``.

Examples
--------
# Real checkpoint, long peptides, nucleus + remask self-correction (recommended default):
python generate_valid.py \
    --ckpt_path checkpoints/td3b.ckpt \
    --length 400 --num_samples 64 \
    --strategy nucleus_remask \
    --device cuda:0 --seed 42 \
    --save_path results/valid_len400.csv

# No checkpoint available -> RANDOM-init model on CPU (development / API smoke test;
# absolute yields are garbage, only the sampling machinery is exercised):
python generate_valid.py --length 200 --num_samples 32 --strategy remask --device cpu

Strategies: baseline, more_steps, top_p (nucleus), top_k, low_temp, remask,
best_of_n, nucleus_remask.  Per-strategy knobs below override the preset defaults.
"""
import argparse
import csv
import logging
import os
import sys

import numpy as np
import torch

ROOT_DIR = os.path.dirname(os.path.abspath(__file__))
if ROOT_DIR not in sys.path:
    sys.path.insert(0, ROOT_DIR)

from sampling_strategies import generate, build_random_model, available_strategies
from utils.app import PeptideAnalyzer

logger = logging.getLogger("generate_valid")
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")


def _load_model(ckpt_path, device, base_path, hidden_size, n_layers, n_heads):
    """Load the real checkpoint via ``inference.load_model`` when available; otherwise
    fall back to a RANDOM-init model (imports of ``inference`` are done lazily so the
    random path has no heavy dependencies)."""
    if ckpt_path and os.path.isfile(ckpt_path):
        logger.info("Loading real checkpoint from %s", ckpt_path)
        from inference import load_model  # reuse the canonical loader
        model, tokenizer = load_model(ckpt_path, device)
        return model, tokenizer, False
    if ckpt_path:
        logger.warning("Checkpoint %s not found -- falling back to RANDOM-init model.", ckpt_path)
    else:
        logger.warning("No --ckpt_path given -- using RANDOM-init model "
                       "(yields are meaningless; API/mechanism check only).")
    model, tokenizer = build_random_model(
        device=device, hidden_size=hidden_size, n_layers=n_layers,
        n_heads=n_heads, base_path=base_path)
    return model, tokenizer, True


def build_parser():
    p = argparse.ArgumentParser(description="Generate valid peptides with a chosen sampling strategy.")
    p.add_argument("--ckpt_path", type=str, default=None,
                   help="Path to TD3B checkpoint. If missing/omitted, a random-init model is used.")
    p.add_argument("--base_path", type=str, default=ROOT_DIR, help="Repo root (for tokenizer files).")
    p.add_argument("--length", type=int, default=200, help="Target sequence length (tokens).")
    p.add_argument("--num_samples", type=int, default=64, help="Number of sequences to generate.")
    p.add_argument("--strategy", type=str, default="nucleus_remask",
                   choices=available_strategies(), help="Sampling strategy.")
    p.add_argument("--device", type=str, default="cuda:0")
    p.add_argument("--seed", type=int, default=42)
    p.add_argument("--save_path", type=str, default=None,
                   help="CSV path to save VALID sequences (default: results/valid_<strategy>_len<L>.csv).")
    # strategy knobs (None -> use the strategy preset default)
    p.add_argument("--num_steps", type=int, default=128, help="Base reverse-diffusion steps.")
    p.add_argument("--eps", type=float, default=1e-5)
    p.add_argument("--temperature", type=float, default=None, help="<1 sharpens logits (low_temp).")
    p.add_argument("--top_p", type=float, default=None, help="Nucleus mass in (0,1].")
    p.add_argument("--top_k", type=int, default=None, help="Top-k tokens per position.")
    p.add_argument("--steps_per_token", type=float, default=None,
                   help="more_steps: num_steps = max(num_steps, round(steps_per_token*length)).")
    p.add_argument("--remask_rounds", type=int, default=None, help="Self-correction rounds.")
    p.add_argument("--remask_frac", type=float, default=None, help="Fraction of lowest-conf tokens to remask.")
    p.add_argument("--remask_steps", type=int, default=None, help="Re-denoise steps per remask round.")
    p.add_argument("--best_of_n", type=int, default=None, help="Oversample N per slot, keep first valid.")
    # random-fallback architecture (ignored when a real checkpoint loads)
    p.add_argument("--hidden_size", type=int, default=768)
    p.add_argument("--n_layers", type=int, default=8)
    p.add_argument("--n_heads", type=int, default=8)
    return p


def main():
    args = build_parser().parse_args()

    torch.manual_seed(args.seed)
    np.random.seed(args.seed)
    device = torch.device(args.device if (args.device.startswith("cpu") or torch.cuda.is_available()) else "cpu")

    model, tokenizer, is_random = _load_model(
        args.ckpt_path, device, args.base_path, args.hidden_size, args.n_layers, args.n_heads)
    analyzer = PeptideAnalyzer()

    logger.info("Generating %d sequences of length %d with strategy=%s on %s",
                args.num_samples, args.length, args.strategy, device)

    tokens, sequences, valid_mask, stats = generate(
        model, tokenizer, analyzer,
        batch_size=args.num_samples, length=args.length, strategy=args.strategy,
        num_steps=args.num_steps, eps=args.eps,
        temperature=args.temperature, top_p=args.top_p, top_k=args.top_k,
        steps_per_token=args.steps_per_token,
        remask_rounds=args.remask_rounds, remask_frac=args.remask_frac,
        remask_steps=args.remask_steps, best_of_n=args.best_of_n,
        verbose=False,
    )

    valid_seqs = [s for s, v in zip(sequences, valid_mask) if v]

    save_path = args.save_path or os.path.join(
        args.base_path, "results", f"valid_{args.strategy}_len{args.length}.csv")
    os.makedirs(os.path.dirname(os.path.abspath(save_path)), exist_ok=True)
    with open(save_path, "w", newline="") as f:
        w = csv.writer(f)
        w.writerow(["idx", "sequence", "n_chars"])
        for i, s in enumerate(valid_seqs):
            w.writerow([i, s, len(s)])

    print("\n" + "=" * 66)
    print(f"  strategy         : {stats['strategy']}")
    print(f"  length           : {stats['length']}")
    print(f"  num_samples      : {stats['batch_size']}")
    print(f"  num_steps        : {stats['num_steps']}"
          f"   (temp={stats['temperature']}, top_p={stats['top_p']}, top_k={stats['top_k']})")
    print(f"  remask           : rounds={stats['remask_rounds']} frac={stats['remask_frac']} "
          f"steps={stats['remask_steps']}   best_of_n={stats['best_of_n']}")
    if len(stats.get("round_valid_counts", [])) > 1:
        print(f"  valid per round  : {stats['round_valid_counts']}  (round 0 = before remask)")
    print(f"  VALID YIELD      : {stats['valid_count']}/{stats['batch_size']} "
          f"= {stats['valid_rate']:.1%}")
    print(f"  wall time        : {stats['wall_time_s']}s")
    print(f"  saved valid seqs : {save_path}  ({len(valid_seqs)} rows)")
    if is_random:
        print("  NOTE: RANDOM-init model -- yields are meaningless; rerun with --ckpt_path "
              "<real.ckpt> for real numbers.")
    print("=" * 66)


if __name__ == "__main__":
    main()