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