File size: 7,252 Bytes
914512c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""bioai CLI dispatcher.

Usage::

    python -m bioai --help
    python -m bioai train --epochs 10 --batch-size 64
    python -m bioai train-vae --epochs 10
    python -m bioai train-pinn --epochs 100
    python -m bioai rank --input candidates.txt --output ranked.csv
    python -m bioai design --user-text "Brown planthopper in rice paddy in Tamil Nadu"
"""

from __future__ import annotations

import argparse
import sys
from typing import List, Optional


def main(argv: Optional[List[str]] = None) -> int:
    parser = argparse.ArgumentParser(
        prog="bioai",
        description="Biopesticide-AI: dsRNA biopesticide design pipeline (AMD ROCm + Fireworks AI).",
    )
    sub = parser.add_subparsers(dest="cmd", required=True)

    # train
    p_train = sub.add_parser("train", help="Train SiRNACNN on the multi-task siRNA dataset.")
    p_train.add_argument("--epochs", type=int, default=10)
    p_train.add_argument("--batch-size", type=int, default=64)
    p_train.add_argument("--data", type=str, default="data/processed/training_data.csv")
    p_train.add_argument("--device", type=str, default="auto", choices=["auto", "cpu", "cuda"])
    p_train.add_argument("--lr", type=float, default=1e-3)
    p_train.add_argument("--patience", type=int, default=5)
    p_train.add_argument("--use-caduceus", action="store_true")
    p_train.add_argument("--checkpoint", type=str,
                         default=None)

    # train-vae
    p_vae = sub.add_parser("train-vae", help="Train DiscreteVAE on 200-nt precursors.")
    p_vae.add_argument("--epochs", type=int, default=10)
    p_vae.add_argument("--batch-size", type=int, default=32)
    p_vae.add_argument("--device", type=str, default="auto", choices=["auto", "cpu", "cuda"])
    p_vae.add_argument("--checkpoint", type=str,
                       default=None)

    # train-pinn
    p_pinn = sub.add_parser("train-pinn", help="Train DegradationPINN on synthetic fate data.")
    p_pinn.add_argument("--epochs", type=int, default=100)
    p_pinn.add_argument("--batch-size", type=int, default=64)
    p_pinn.add_argument("--device", type=str, default="auto", choices=["auto", "cpu", "cuda"])
    p_pinn.add_argument("--n-samples", type=int, default=1024)
    p_pinn.add_argument("--checkpoint", type=str,
                        default=None)

    # rank
    p_rank = sub.add_parser("rank", help="Rank dsRNA candidates.")
    p_rank.add_argument("--input", type=str, required=True)
    p_rank.add_argument("--output", type=str, default="ranked.csv")
    p_rank.add_argument("--device", type=str, default="auto", choices=["auto", "cpu", "cuda"])
    p_rank.add_argument("--top-k", type=int, default=20)
    p_rank.add_argument("--sirna-checkpoint", type=str,
                        default=None)
    p_rank.add_argument("--pinn-checkpoint", type=str,
                        default=None)

    # design
    p_design = sub.add_parser("design", help="Run the end-to-end design pipeline.")
    p_design.add_argument("--user-text", type=str, required=True)
    p_design.add_argument("--device", type=str, default="auto", choices=["auto", "cpu", "cuda"])
    p_design.add_argument("--pest-fasta", type=str, default=None)
    p_design.add_argument("--safety-fasta", type=str, default=None)
    p_design.add_argument("--sirna-checkpoint", type=str, default=None)
    p_design.add_argument("--pinn-checkpoint", type=str, default=None)
    p_design.add_argument("--top-k", type=int, default=10)
    p_design.add_argument("--max-transcripts", type=int, default=5)
    p_design.add_argument("--pest-species", type=str, default=None,
                          help="skip LLM parsing and use this species directly (e.g. nilaparvata_lugens)")

    # web (replaces the old Gradio `ui` subcommand)
    p_web = sub.add_parser("web", help="Launch the FastAPI web UI.")
    p_web.add_argument("--host", type=str, default="0.0.0.0")
    p_web.add_argument("--port", type=int, default=7860)
    p_web.add_argument("--reload", action="store_true", help="auto-reload on file changes (dev mode)")

    args = parser.parse_args(argv)

    if args.cmd == "train":
        from .training.train import train as _train
        from .paths import SIRNA_CHECKPOINT
        ckpt = __import__("pathlib").Path(args.checkpoint) if args.checkpoint else SIRNA_CHECKPOINT
        _train(
            csv_path=args.data,
            epochs=args.epochs,
            batch_size=args.batch_size,
            lr=args.lr,
            patience=args.patience,
            device=args.device,
            use_caduceus=args.use_caduceus,
            checkpoint_path=ckpt,
        )
        return 0

    if args.cmd == "train-vae":
        from .training.train_vae import train as _train_vae
        from .paths import VAE_CHECKPOINT
        ckpt = __import__("pathlib").Path(args.checkpoint) if args.checkpoint else VAE_CHECKPOINT
        _train_vae(
            epochs=args.epochs,
            batch_size=args.batch_size,
            device=args.device,
            checkpoint_path=ckpt,
        )
        return 0

    if args.cmd == "train-pinn":
        from .training.train_pinn import train as _train_pinn
        from .paths import PINN_CHECKPOINT
        ckpt = __import__("pathlib").Path(args.checkpoint) if args.checkpoint else PINN_CHECKPOINT
        _train_pinn(
            epochs=args.epochs,
            batch_size=args.batch_size,
            device=args.device,
            n_samples=args.n_samples,
            checkpoint_path=ckpt,
        )
        return 0

    if args.cmd == "rank":
        from .inference.ranker import main as _rank_main
        from .paths import SIRNA_CHECKPOINT, PINN_CHECKPOINT
        sirna_ckpt = args.sirna_checkpoint or str(SIRNA_CHECKPOINT)
        pinn_ckpt = args.pinn_checkpoint or str(PINN_CHECKPOINT)
        return _rank_main([
            "--input", args.input,
            "--output", args.output,
            "--device", args.device,
            "--top-k", str(args.top_k),
            "--sirna-checkpoint", sirna_ckpt,
            "--pinn-checkpoint", pinn_ckpt,
        ])

    if args.cmd == "design":
        from .orchestrator import main as _design_main
        from .paths import DEFAULT_PEST_FASTA, DEFAULT_SAFETY_FASTA, SIRNA_CHECKPOINT, PINN_CHECKPOINT
        pest = args.pest_fasta or str(DEFAULT_PEST_FASTA)
        safety = args.safety_fasta or str(DEFAULT_SAFETY_FASTA)
        sirna_ckpt = args.sirna_checkpoint or str(SIRNA_CHECKPOINT)
        pinn_ckpt = args.pinn_checkpoint or str(PINN_CHECKPOINT)
        return _design_main([
            "--user-text", args.user_text,
            "--device", args.device,
            "--pest-fasta", pest,
            "--safety-fasta", safety,
            "--sirna-checkpoint", sirna_ckpt,
            "--pinn-checkpoint", pinn_ckpt,
            "--top-k", str(args.top_k),
            "--max-transcripts", str(args.max_transcripts),
        ])

    if args.cmd == "web":
        from .web.api import main as _web_main
        new_argv = ["bioai-web", "--host", args.host, "--port", str(args.port)]
        if args.reload:
            new_argv.append("--reload")
        sys.argv = new_argv
        _web_main()
        return 0

    parser.print_help()
    return 1


if __name__ == "__main__":
    sys.exit(main())