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