Spaces:
Sleeping
Sleeping
| """bioai.training.train -- unified training script for the siRNA CNN. | |
| Trains :class:`bioai.models.SiRNACNN` (or :class:`CaduceusAdapter` if Caduceus | |
| is installed) on the multi-task objective | |
| loss = 1.0 * BCE(efficacy) + 0.5 * BCE(safety_per_species) | |
| with Adam (lr=1e-3) and early stopping (patience=5) on validation F1. | |
| CLI:: | |
| python -m bioai.training.train --epochs 10 --batch-size 64 \ | |
| --data data/processed/training_data.csv --device auto | |
| The CSV is the one produced by Task 3-B's ``build_training_csv.py`` and must | |
| contain at least these columns:: | |
| sirna_seq, knockdown_pct, offtarget_<species>... | |
| ``sirna_seq`` is 21 nt; ``knockdown_pct`` is the regression target rescaled | |
| to [0,1]; the ``offtarget_*`` columns are binarised at the 0.3 threshold and | |
| used as the multi-label safety target. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import sys | |
| from pathlib import Path | |
| from typing import List, Tuple | |
| import numpy as np | |
| import pandas as pd | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from sklearn.metrics import f1_score, roc_auc_score | |
| from torch.utils.data import DataLoader, Dataset, random_split | |
| from ..models.caduceus_adapter import CaduceusAdapter | |
| from ..models.sirna_cnn import SiRNACNN, resolve_device | |
| from ..sequence_utils import encode_batch, SAFETY_SPECIES | |
| # Default safety-species columns produced by build_training_csv.py | |
| SAFETY_COLS = [f"offtarget_{sp}" for sp in SAFETY_SPECIES] | |
| # Portable checkpoint path (resolved from bioai.paths) | |
| from bioai.paths import SIRNA_CHECKPOINT as CHECKPOINT_PATH # noqa: E402 | |
| # --------------------------------------------------------------------------- # | |
| # Dataset | |
| # --------------------------------------------------------------------------- # | |
| class SirnaDataset(Dataset): | |
| def __init__(self, csv_path: str | Path, seq_len: int = 21, | |
| safety_threshold: float = 0.3, | |
| safety_cols: List[str] | None = None): | |
| self.df = pd.read_csv(csv_path) | |
| self.seq_len = seq_len | |
| self.safety_threshold = safety_threshold | |
| self.safety_cols = safety_cols or [ | |
| c for c in self.df.columns if c.startswith("offtarget_") | |
| ] | |
| # Some CSVs name the sequence column "sequence", some "sirna_seq". | |
| self.seq_col = "sirna_seq" if "sirna_seq" in self.df.columns else "sequence" | |
| # Some CSVs use knockdown_pct (regression), some use pest_label (binary). | |
| # We always train on knockdown_pct scaled to [0,1] if present; otherwise | |
| # fall back to pest_label. | |
| if "knockdown_pct" in self.df.columns: | |
| self.eff_col = "knockdown_pct" | |
| else: | |
| self.eff_col = "pest_label" | |
| def __len__(self) -> int: | |
| return len(self.df) | |
| def __getitem__(self, idx: int): | |
| row = self.df.iloc[idx] | |
| seq = str(row[self.seq_col]) | |
| x = torch.tensor(encode_batch([seq], max_len=self.seq_len)[0], dtype=torch.float32) | |
| # efficacy: knockdown_pct is in [0,1] already; pest_label is 0/1. | |
| eff = torch.tensor(float(row[self.eff_col]), dtype=torch.float32) | |
| # safety: binarise offtarget scores at the threshold. | |
| safety = torch.tensor( | |
| [float(row[c] > self.safety_threshold) for c in self.safety_cols], | |
| dtype=torch.float32, | |
| ) | |
| return x, eff, safety | |
| # --------------------------------------------------------------------------- # | |
| # Training loop | |
| # --------------------------------------------------------------------------- # | |
| def train( | |
| csv_path: str | Path, | |
| epochs: int = 10, | |
| batch_size: int = 64, | |
| lr: float = 1e-3, | |
| patience: int = 5, | |
| device: str = "auto", | |
| use_caduceus: bool = False, | |
| checkpoint_path: Path | None = None, | |
| val_frac: float = 0.15, | |
| safety_weight: float = 0.5, | |
| safety_cols: List[str] | None = None, | |
| ) -> str: | |
| """Run training. Returns the path to the saved checkpoint.""" | |
| device_t = resolve_device(device) | |
| print(f"[train] device = {device_t}") | |
| checkpoint_path = checkpoint_path or CHECKPOINT_PATH | |
| checkpoint_path.parent.mkdir(parents=True, exist_ok=True) | |
| # ----- model ----------------------------------------------------------- | |
| seq_len = 21 | |
| num_safety = len(safety_cols) if safety_cols else len(SAFETY_COLS) | |
| if use_caduceus: | |
| model = CaduceusAdapter(seq_len=seq_len, num_safety_species=num_safety, device=device) | |
| print("[train] using CaduceusAdapter (backend may fall back to CNN)") | |
| else: | |
| model = SiRNACNN(seq_len=seq_len, num_safety_species=num_safety) | |
| print("[train] using SiRNACNN") | |
| model.to(device_t) | |
| # ----- data ------------------------------------------------------------ | |
| full = SirnaDataset(csv_path, seq_len=seq_len, safety_cols=safety_cols) | |
| n_val = max(1, int(len(full) * val_frac)) | |
| n_train = len(full) - n_val | |
| train_ds, val_ds = random_split( | |
| full, [n_train, n_val], | |
| generator=torch.Generator().manual_seed(42), | |
| ) | |
| train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, drop_last=False) | |
| val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False) | |
| print(f"[train] {n_train} train / {n_val} val samples") | |
| # ----- optimiser / loss ------------------------------------------------ | |
| optimizer = torch.optim.Adam(model.parameters(), lr=lr) | |
| bce = nn.BCELoss() | |
| best_f1 = -1.0 | |
| best_auc = float("nan") | |
| patience_counter = 0 | |
| for epoch in range(1, epochs + 1): | |
| # --- train --------------------------------------------------------- | |
| model.train() | |
| train_loss = 0.0 | |
| n_batches = 0 | |
| for x, eff, safety in train_loader: | |
| x = x.to(device_t) | |
| eff = eff.to(device_t).float().view(-1, 1) | |
| safety = safety.to(device_t).float() | |
| optimizer.zero_grad() | |
| eff_pred, safety_pred = model(x) | |
| # safety_pred may have a different number of columns if the user | |
| # passed a custom safety_cols list that doesn't match the model's | |
| # num_safety_species. Clip to the smaller of the two. | |
| min_s = min(safety_pred.size(1), safety.size(1)) | |
| loss_eff = F.binary_cross_entropy(eff_pred.clamp(1e-6, 1 - 1e-6), eff) | |
| loss_safe = bce(safety_pred[:, :min_s].clamp(1e-6, 1 - 1e-6), safety[:, :min_s]) | |
| loss = 1.0 * loss_eff + safety_weight * loss_safe | |
| loss.backward() | |
| optimizer.step() | |
| train_loss += loss.item() | |
| n_batches += 1 | |
| train_loss /= max(1, n_batches) | |
| # --- validate ------------------------------------------------------ | |
| model.eval() | |
| all_true: List[float] = [] | |
| all_pred: List[float] = [] | |
| val_loss = 0.0 | |
| n_val_batches = 0 | |
| with torch.no_grad(): | |
| for x, eff, safety in val_loader: | |
| x = x.to(device_t) | |
| eff = eff.to(device_t).float().view(-1, 1) | |
| safety = safety.to(device_t).float() | |
| eff_pred, safety_pred = model(x) | |
| min_s = min(safety_pred.size(1), safety.size(1)) | |
| loss_eff = F.binary_cross_entropy(eff_pred.clamp(1e-6, 1 - 1e-6), eff) | |
| loss_safe = bce(safety_pred[:, :min_s].clamp(1e-6, 1 - 1e-6), safety[:, :min_s]) | |
| val_loss += (1.0 * loss_eff + safety_weight * loss_safe).item() | |
| n_val_batches += 1 | |
| all_true.extend(eff.cpu().numpy().reshape(-1).tolist()) | |
| all_pred.extend(eff_pred.cpu().numpy().reshape(-1).tolist()) | |
| val_loss /= max(1, n_val_batches) | |
| # AUC / F1 on the binary "efficacious" label (threshold 0.5 on both | |
| # ground truth and prediction; treat knockdown_pct >= 0.5 as positive). | |
| y_true = [1 if t >= 0.5 else 0 for t in all_true] | |
| y_pred_bin = [1 if p >= 0.5 else 0 for p in all_pred] | |
| try: | |
| auc = roc_auc_score(y_true, all_pred) if len(set(y_true)) > 1 else float("nan") | |
| except Exception: | |
| auc = float("nan") | |
| f1 = f1_score(y_true, y_pred_bin, zero_division=0) | |
| print( | |
| f"Epoch {epoch:3d}/{epochs}: " | |
| f"train_loss={train_loss:.4f} val_loss={val_loss:.4f} " | |
| f"val_auc={auc:.4f} val_f1={f1:.4f}" | |
| ) | |
| # --- early stopping ----------------------------------------------- | |
| if f1 > best_f1: | |
| best_f1 = f1 | |
| best_auc = auc | |
| patience_counter = 0 | |
| torch.save(model.state_dict(), checkpoint_path) | |
| print(f" -> saved checkpoint to {checkpoint_path}") | |
| else: | |
| patience_counter += 1 | |
| if patience_counter >= patience: | |
| print(f"[train] early stopping at epoch {epoch} (best F1={best_f1:.4f})") | |
| break | |
| print(f"[train] done. best_val_f1={best_f1:.4f} best_val_auc={best_auc:.4f}") | |
| return str(checkpoint_path) | |
| # --------------------------------------------------------------------------- # | |
| # CLI | |
| # --------------------------------------------------------------------------- # | |
| def main(argv: List[str] | None = None) -> int: | |
| p = argparse.ArgumentParser(description="Train SiRNACNN on the multi-task siRNA dataset.") | |
| p.add_argument("--epochs", type=int, default=10) | |
| p.add_argument("--batch-size", type=int, default=64) | |
| p.add_argument("--data", type=str, default="data/processed/training_data.csv") | |
| p.add_argument("--device", type=str, default="auto", choices=["auto", "cpu", "cuda"]) | |
| p.add_argument("--lr", type=float, default=1e-3) | |
| p.add_argument("--patience", type=int, default=5) | |
| p.add_argument("--use-caduceus", action="store_true", | |
| help="Use CaduceusAdapter (loads Caduceus if available, else CNN).") | |
| p.add_argument("--checkpoint", type=str, default=str(CHECKPOINT_PATH)) | |
| args = p.parse_args(argv) | |
| ckpt = 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=Path(args.checkpoint), | |
| ) | |
| print(f"[train] checkpoint: {ckpt}") | |
| return 0 | |
| if __name__ == "__main__": | |
| sys.exit(main()) | |