Biopesticide-AI / bioai /training /train_pinn.py
flvcko's picture
Biopesticide-AI: AMD Hackathon Unicorn Track submission
914512c
Raw
History Blame Contribute Delete
9.04 kB
"""bioai.training.train_pinn -- train the DegradationPINN on synthetic fate data.
We do not have real environmental fate measurements for dsRNA (collecting them
requires a multi-week field trial with LC-MS). For the hackathon we therefore
generate SYNTHETIC (features, half-life) pairs from a hand-coded physical
heuristic -- higher temperature, UV, salinity, and humidity all increase the
degradation rate ``k``; higher GC content and length slightly stabilise the
duplex -- and use them to teach the PINN the right qualitative behaviour.
The physics-consistency loss on ``C(t) = C0 * exp(-k * t)`` then ensures the
network respects the ODE even at feature combos it has never seen.
The PINN is consumed by the ranker to penalise candidates with a half-life
under 6 hours, so the synthetic training just needs to produce a network
that says "too short" for hot/sunny/wet conditions and "long enough" for
cool/dry conditions. The MSE + physics loss combo does exactly that.
CLI::
python -m bioai.training.train_pinn --epochs 100
"""
from __future__ import annotations
import argparse
import math
import sys
from pathlib import Path
from typing import List
import numpy as np
import torch
import torch.nn as nn
from ..models.pinn_fate import DegradationPINN, PINN_FEATURE_NAMES
from ..models.sirna_cnn import resolve_device
# Portable checkpoint path (resolved from bioai.paths)
from bioai.paths import PINN_CHECKPOINT as CHECKPOINT_PATH # noqa: E402
# --------------------------------------------------------------------------- #
# Synthetic fate generator
# --------------------------------------------------------------------------- #
def generate_synthetic_fate_data(
n_samples: int = 1024,
seed: int = 13,
) -> tuple[np.ndarray, np.ndarray]:
"""Generate ``(features, half_life_hours)`` pairs.
Features (8-dim) follow realistic ranges; the half-life is derived from
a hand-coded rate that captures the qualitative physics (Arrhenius-style
temperature dependence, UV photocatalysis, salinity-driven hydrolysis,
GC-stabilisation, length-stabilisation). This is a *teaching signal*,
not a measurement -- see module docstring for the rationale.
"""
rng = np.random.default_rng(seed)
# Realistic ranges per feature
temp = rng.uniform(10.0, 40.0, n_samples) # Celsius
pH = rng.uniform(5.0, 9.0, n_samples)
uv = rng.uniform(0.0, 12.0, n_samples) # UV index
gc = rng.uniform(0.3, 0.7, n_samples) # fraction
length = rng.uniform(50.0, 500.0, n_samples) # nt
sal = rng.uniform(0.0, 35.0, n_samples) # ppt
clay = rng.uniform(0.0, 60.0, n_samples) # %
hum = rng.uniform(10.0, 100.0, n_samples) # %
features = np.stack([temp, pH, uv, gc, length, sal, clay, hum], axis=1).astype(np.float32)
# Hand-coded rate (1/hours). Each term contributes multiplicatively.
# Reference: dsRNA in soil literature reports half-lives of 1-72 hours
# depending on conditions; we target that range.
arrhenius = np.exp((temp - 25.0) / 12.0) # Q10-style
uv_factor = 1.0 + 0.15 * uv # UV photocatalysis
sal_factor = 1.0 + 0.03 * sal # salinity hydrolysis
hum_factor = 1.0 + 0.005 * (hum - 50.0) # humidity mild effect
gc_stabiliser = 1.0 / (0.5 + gc) # high GC -> slower
len_stabiliser = 200.0 / length # long duplex -> slower
base_rate = 0.10 # 1/hours at reference
k = base_rate * arrhenius * uv_factor * sal_factor * hum_factor * gc_stabiliser * len_stabiliser
# Add small noise so the PINN can't just memorise the heuristic.
k = k * rng.uniform(0.9, 1.1, n_samples)
half_life = np.log(2.0) / k
return features, half_life.astype(np.float32)
# --------------------------------------------------------------------------- #
# Training
# --------------------------------------------------------------------------- #
def train(
epochs: int = 100,
batch_size: int = 64,
lr: float = 1e-3,
device: str = "auto",
n_samples: int = 1024,
physics_weight: float = 0.5,
checkpoint_path: Path | None = None,
) -> str:
device_t = resolve_device(device)
print(f"[train_pinn] device = {device_t}")
features_np, hl_np = generate_synthetic_fate_data(n_samples=n_samples)
# Convert half-life to rate (the PINN output) for the supervised MSE.
k_np = (np.log(2.0) / hl_np).astype(np.float32)
# Hold out 20% for validation.
n_val = max(1, int(0.2 * len(features_np)))
rng = np.random.default_rng(42)
perm = rng.permutation(len(features_np))
val_idx, train_idx = perm[:n_val], perm[n_val:]
feat_tr = torch.tensor(features_np[train_idx], dtype=torch.float32, device=device_t)
k_tr = torch.tensor(k_np[train_idx], dtype=torch.float32, device=device_t).view(-1, 1)
feat_val = torch.tensor(features_np[val_idx], dtype=torch.float32, device=device_t)
k_val = torch.tensor(k_np[val_idx], dtype=torch.float32, device=device_t).view(-1, 1)
print(f"[train_pinn] {len(train_idx)} train / {len(val_idx)} val synthetic samples")
model = DegradationPINN(feature_dim=8).to(device_t)
optimizer = torch.optim.Adam(model.parameters(), lr=lr)
checkpoint_path = checkpoint_path or CHECKPOINT_PATH
checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
# Time grid for physics-consistency loss (hours).
t_grid = torch.linspace(0.0, 24.0, 13, device=device_t) # every 2 hours
C0 = 1.0
best_val_mse = float("inf")
for epoch in range(1, epochs + 1):
model.train()
# Mini-batch gradient descent over the training set.
perm_t = torch.randperm(len(feat_tr), device=device_t)
total_supervised = 0.0
total_physics = 0.0
n_batches = 0
for i in range(0, len(feat_tr), batch_size):
idx = perm_t[i:i + batch_size]
f_b = feat_tr[idx]
k_b = k_tr[idx]
optimizer.zero_grad()
# Supervised MSE on rate.
k_pred = model.predict_rate(f_b)
loss_sup = nn.functional.mse_loss(k_pred, k_b)
# Physics consistency: generate a target trajectory from the
# GROUND-TRUTH rate and ask the PINN to reproduce it from
# features alone. This forces k_pred to match k_b *via* the ODE.
C_target = C0 * torch.exp(-k_b * t_grid.unsqueeze(0)) # (B, T)
C_pred = model.predict_concentration(C0, t_grid, f_b) # (B, T)
loss_phys = nn.functional.mse_loss(C_pred, C_target)
loss = loss_sup + physics_weight * loss_phys
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
total_supervised += loss_sup.item()
total_physics += loss_phys.item()
n_batches += 1
# Validation
model.eval()
with torch.no_grad():
k_val_pred = model.predict_rate(feat_val)
val_mse = nn.functional.mse_loss(k_val_pred, k_val).item()
hl_pred = model.half_life(feat_val).cpu().numpy().reshape(-1)
hl_true = (math.log(2.0) / k_val.cpu().numpy().reshape(-1))
if epoch % 10 == 0 or epoch == 1:
print(
f"Epoch {epoch:3d}/{epochs}: "
f"sup_loss={total_supervised / max(1, n_batches):.6f} "
f"phys_loss={total_physics / max(1, n_batches):.6f} "
f"val_mse_k={val_mse:.6f} "
f"val_hl_mean_pred={hl_pred.mean():.2f}h "
f"val_hl_mean_true={hl_true.mean():.2f}h"
)
if val_mse < best_val_mse:
best_val_mse = val_mse
torch.save(model.state_dict(), checkpoint_path)
print(f"[train_pinn] done. best_val_mse_k={best_val_mse:.6f}")
print(f"[train_pinn] checkpoint: {checkpoint_path}")
return str(checkpoint_path)
# --------------------------------------------------------------------------- #
# CLI
# --------------------------------------------------------------------------- #
def main(argv: List[str] | None = None) -> int:
p = argparse.ArgumentParser(description="Train the DegradationPINN on synthetic fate data.")
p.add_argument("--epochs", type=int, default=100)
p.add_argument("--batch-size", type=int, default=64)
p.add_argument("--lr", type=float, default=1e-3)
p.add_argument("--device", type=str, default="auto", choices=["auto", "cpu", "cuda"])
p.add_argument("--n-samples", type=int, default=1024)
p.add_argument("--checkpoint", type=str, default=str(CHECKPOINT_PATH))
args = p.parse_args(argv)
train(
epochs=args.epochs,
batch_size=args.batch_size,
lr=args.lr,
device=args.device,
n_samples=args.n_samples,
checkpoint_path=Path(args.checkpoint),
)
return 0
if __name__ == "__main__":
sys.exit(main())