"""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())