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