Spaces:
Sleeping
Sleeping
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())
|