Spaces:
Sleeping
Sleeping
File size: 7,165 Bytes
914512c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 | """bioai.training.train_vae -- train the DiscreteVAE on 200-nt dsRNA precursors.
CLI::
python -m bioai.training.train_vae --epochs 10
Extracts 200-nt precursors from ``data/synthetic/pest_transcripts.fasta`` (or
real data if available at ``data/external/pest_transcripts.fasta``), tiles them
into 200-nt windows with 50% overlap, and trains the VAE for the given number
of epochs. Saves the checkpoint to ``checkpoints/vae_best.pt`` under the project root.
This is a STRETCH GOAL: even an untrained VAE is enough for the demo (we just
need the ``sample()`` method to produce *some* 200-nt precursor that the
downstream ranker can dice into siRNAs). But training does make the samples
look more like real transcripts (GC content, k-mer frequencies).
"""
from __future__ import annotations
import argparse
import sys
from pathlib import Path
from typing import List, Tuple
import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader, Dataset
from ..models.discrete_vae import DiscreteVAE
from ..models.sirna_cnn import resolve_device
from ..sequence_utils import fasta_iter, tile_sequence
# Portable paths (resolved from bioai.paths)
from bioai.paths import VAE_CHECKPOINT as CHECKPOINT_PATH # noqa: E402
from bioai.paths import EXTERNAL_DIR, SYNTHETIC_DIR # noqa: E402
# A/C/G/T -> {0,1,2,3}
BASE_TO_IDX = {"A": 0, "C": 1, "G": 2, "T": 3, "U": 3}
# --------------------------------------------------------------------------- #
# Dataset
# --------------------------------------------------------------------------- #
class PrecursorDataset(Dataset):
def __init__(self, fasta_paths: List[Path], seq_len: int = 200, step: int = 100,
max_per_transcript: int = 5, max_total: int = 2000):
self.seq_len = seq_len
self.records: List[torch.Tensor] = []
for path in fasta_paths:
if not path.exists():
continue
for _, seq in fasta_iter(path):
# skip overly short transcripts
if len(seq) < seq_len:
continue
windows = tile_sequence(seq, window=seq_len, step=step,
max_candidates=max_per_transcript)
for _, _, sub in windows:
if len(self.records) >= max_total:
break
toks = [BASE_TO_IDX.get(b, 0) for b in sub[:seq_len]]
if len(toks) < seq_len:
# pad with A
toks = toks + [0] * (seq_len - len(toks))
self.records.append(torch.tensor(toks, dtype=torch.long))
if len(self.records) >= max_total:
break
if len(self.records) >= max_total:
break
if not self.records:
# fallback: generate random sequences so training still runs
print("[train_vae] no real transcripts found -- using random 200-mers")
for _ in range(256):
self.records.append(torch.randint(0, 4, (seq_len,)))
def __len__(self) -> int:
return len(self.records)
def __getitem__(self, idx: int) -> torch.Tensor:
return self.records[idx]
# --------------------------------------------------------------------------- #
# Training
# --------------------------------------------------------------------------- #
def train(
epochs: int = 10,
batch_size: int = 32,
lr: float = 1e-3,
device: str = "auto",
seq_len: int = 200,
fasta_paths: List[Path] | None = None,
checkpoint_path: Path | None = None,
) -> str:
device_t = resolve_device(device)
print(f"[train_vae] device = {device_t}")
if fasta_paths is None:
candidates = [
EXTERNAL_DIR / "pest_transcripts.fasta",
SYNTHETIC_DIR / "pest_transcripts.fasta",
]
fasta_paths = candidates
print(f"[train_vae] FASTA sources: {[str(p) for p in fasta_paths]}")
ds = PrecursorDataset(fasta_paths, seq_len=seq_len)
print(f"[train_vae] {len(ds)} precursors")
loader = DataLoader(ds, batch_size=batch_size, shuffle=True, drop_last=False)
model = DiscreteVAE(
seq_len=seq_len, vocab_size=4, latent_dim=64,
hidden_dim=128, num_heads=4, num_layers=2,
).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)
best_loss = float("inf")
for epoch in range(1, epochs + 1):
model.train()
running_total = 0.0
running_recon = 0.0
running_kl = 0.0
n_batches = 0
for tokens in loader:
tokens = tokens.to(device_t)
optimizer.zero_grad()
recon_logits, mu, logvar = model(tokens)
loss, info = DiscreteVAE.elbo_loss(recon_logits, tokens, mu, logvar, beta=1.0)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
running_total += info["total"]
running_recon += info["recon"]
running_kl += info["kl"]
n_batches += 1
avg_total = running_total / max(1, n_batches)
avg_recon = running_recon / max(1, n_batches)
avg_kl = running_kl / max(1, n_batches)
print(
f"Epoch {epoch:3d}/{epochs}: "
f"loss={avg_total:.4f} recon={avg_recon:.4f} kl={avg_kl:.4f}"
)
if avg_total < best_loss:
best_loss = avg_total
torch.save(model.state_dict(), checkpoint_path)
print(f" -> saved checkpoint to {checkpoint_path}")
# quick sanity: sample one precursor and check GC content
model.eval()
with torch.no_grad():
sample = model.sample(1, device=device_t).cpu().numpy()[0]
gc = (sum(1 for t in sample if t in (1, 2)) / len(sample)) * 100
print(f"[train_vae] sample GC = {gc:.1f}% (target ~50%)")
print(f"[train_vae] done. best_loss={best_loss:.4f}")
return str(checkpoint_path)
# --------------------------------------------------------------------------- #
# CLI
# --------------------------------------------------------------------------- #
def main(argv: List[str] | None = None) -> int:
p = argparse.ArgumentParser(description="Train the DiscreteVAE on 200-nt precursors.")
p.add_argument("--epochs", type=int, default=10)
p.add_argument("--batch-size", type=int, default=32)
p.add_argument("--lr", type=float, default=1e-3)
p.add_argument("--device", type=str, default="auto", choices=["auto", "cpu", "cuda"])
p.add_argument("--seq-len", type=int, default=200)
p.add_argument("--checkpoint", type=str, default=str(CHECKPOINT_PATH))
args = p.parse_args(argv)
ckpt = train(
epochs=args.epochs,
batch_size=args.batch_size,
lr=args.lr,
device=args.device,
seq_len=args.seq_len,
checkpoint_path=Path(args.checkpoint),
)
print(f"[train_vae] checkpoint: {ckpt}")
return 0
if __name__ == "__main__":
sys.exit(main())
|