Biopesticide-AI / bioai /__main__.py
flvcko's picture
Biopesticide-AI: AMD Hackathon Unicorn Track submission
914512c
Raw
History Blame Contribute Delete
7.25 kB
"""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())