"""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_... ``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())