JetLaggedByData/scifi-forge
0
1"""2v1_baseline/train.py3Training script — PyTorch port of original TF coursework.4Hyperparameters unchanged: same sequence length, batch size, learning rate.5 6Run:7 python v1_baseline/train.py8 python v1_baseline/train.py --resume # continue from last checkpoint9"""10 11import os12os.environ["TF_CPP_MIN_LOG_LEVEL"] = "3"13os.environ["TF_ENABLE_ONEDNN_OPTS"] = "0"14 15import gc16import argparse17import numpy as np18import torch19import torch.nn as nn20from pathlib import Path21from torch.utils.data import Dataset, DataLoader22 23from lstm_model import build_lstm, loss_fn, VOCAB_SIZE, EMBEDDING_DIM, RNN_UNITS24 25 26# ── Hyperparameters (original notebook values — do not change) ────────────27SEQUENCE_LENGTH = 8028BATCH_SIZE = 12829LEARNING_RATE = 0.00130EPOCHS = 1531MID_EPOCH_SAVE_N = 50_000 # save a mid-epoch checkpoint every N batches32DATA_PATH = Path("../data/raw/internet_archive_scifi_v3.txt")33CHECKPOINT_DIR = Path("checkpoints/lstm_checkpoints")34DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")35 36 37# ── Dataset ───────────────────────────────────────────────────────────────38 39class CharDataset(Dataset):40 """Sliding-window character sequence dataset."""41 42 def __init__(self, int_text: np.ndarray, seq_len: int) -> None:43 self.data = torch.tensor(int_text, dtype=torch.long)44 self.seq_len = seq_len45 46 def __len__(self) -> int:47 return len(self.data) - self.seq_len48 49 def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor]:50 chunk = self.data[idx: idx + self.seq_len + 1]51 return chunk[:-1], chunk[1:] # input, target52 53 54# ── Data loading ──────────────────────────────────────────────────────────55 56def load_and_preprocess(57 data_path: Path = DATA_PATH,58) -> tuple[np.ndarray, list, dict, np.ndarray]:59 """60 Load corpus, strip header, collapse double spaces.61 Identical preprocessing to original notebook.62 """63 text = data_path.read_text(encoding="utf-8")64 text = text[580:149_322_961]65 text = text.replace(" ", " ")66 vocab = sorted(set(text))67 chartoindex = {v: i for i, v in enumerate(vocab)}68 indextochar = np.array(vocab)69 int_text = np.array([chartoindex[c] for c in text])70 71 print(f"Characters: {len(text):,} | Unique: {len(vocab)}")72 return int_text, vocab, chartoindex, indextochar73 74 75# ── Training ──────────────────────────────────────────────────────────────76 77def train(resume: bool = False) -> None:78 """Full training run with early stopping and checkpoint saving."""79 print(f"Device: {DEVICE}")80 81 int_text, vocab, _, _ = load_and_preprocess()82 dataset = CharDataset(int_text, SEQUENCE_LENGTH)83 dataloader = DataLoader(84 dataset,85 batch_size=BATCH_SIZE,86 shuffle=True,87 drop_last=True,88 num_workers=2,89 pin_memory=(DEVICE.type == "cuda"),90 )91 92 model = build_lstm(vocab_size=len(vocab)).to(DEVICE)93 optimiser = torch.optim.Adam(model.parameters(), lr=LEARNING_RATE)94 scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(95 optimiser, patience=1, factor=0.596 )97 scaler = torch.amp.GradScaler("cuda", enabled=(DEVICE.type == "cuda"))98 99 CHECKPOINT_DIR.mkdir(parents=True, exist_ok=True)100 start_epoch = 1101 best_loss = float("inf")102 patience_ctr = 0103 PATIENCE = 2 # matches original EarlyStopping(patience=2)104 105 # Resume from latest checkpoint if requested106 resume_batch = 0 # batch index to skip to when resuming mid-epoch107 if resume:108 mid = CHECKPOINT_DIR / "checkpt_mid.pt"109 epoch_ckpts = sorted(CHECKPOINT_DIR.glob("checkpt_[0-9]*.pt"))110 if mid.exists():111 state = torch.load(mid, map_location=DEVICE, weights_only=False)112 model.load_state_dict(state["model"])113 optimiser.load_state_dict(state["optimiser"])114 if "scaler" in state:115 scaler.load_state_dict(state["scaler"])116 start_epoch = state["epoch"]117 resume_batch = state["batch_idx"] + 1118 best_loss = state["best_loss"]119 print(f"Resumed from mid-epoch checkpoint "120 f"(epoch {start_epoch}, batch {state['batch_idx']})")121 elif epoch_ckpts:122 latest = epoch_ckpts[-1]123 state = torch.load(latest, map_location=DEVICE, weights_only=False)124 model.load_state_dict(state["model"])125 optimiser.load_state_dict(state["optimiser"])126 if "scaler" in state:127 scaler.load_state_dict(state["scaler"])128 start_epoch = state["epoch"] + 1129 best_loss = state["best_loss"]130 print(f"Resumed from {latest} (epoch {state['epoch']})")131 132 for epoch in range(start_epoch, EPOCHS + 1):133 model.train()134 epoch_loss, n_batches = 0.0, 0135 136 for batch_idx, (inputs, targets) in enumerate(dataloader):137 # Skip batches already processed when resuming mid-epoch138 if epoch == start_epoch and batch_idx < resume_batch:139 continue140 resume_batch = 0 # only skip on the first (resumed) epoch141 142 inputs = inputs.to(DEVICE)143 targets = targets.to(DEVICE)144 145 model.reset_states()146 optimiser.zero_grad()147 148 with torch.amp.autocast("cuda", enabled=(DEVICE.type == "cuda")):149 logits, _ = model(inputs)150 # Flatten for cross_entropy: (B*T, vocab) vs (B*T,)151 loss = loss_fn(152 logits.view(-1, logits.size(-1)),153 targets.view(-1),154 )155 scaler.scale(loss).backward()156 scaler.unscale_(optimiser)157 nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)158 scaler.step(optimiser)159 scaler.update()160 161 epoch_loss += loss.item()162 n_batches += 1163 164 if batch_idx % 500 == 0:165 print(f" Epoch {epoch} | batch {batch_idx}/{len(dataloader)} "166 f"| loss {loss.item():.4f}")167 168 # Mid-epoch checkpoint — survives crashes within an epoch169 if batch_idx > 0 and batch_idx % MID_EPOCH_SAVE_N == 0:170 torch.save({171 "epoch": epoch,172 "batch_idx": batch_idx,173 "model": model.state_dict(),174 "optimiser": optimiser.state_dict(),175 "scaler": scaler.state_dict(),176 "best_loss": best_loss,177 }, CHECKPOINT_DIR / "checkpt_mid.pt")178 179 avg_loss = epoch_loss / n_batches180 scheduler.step(avg_loss)181 print(f"Epoch {epoch}/{EPOCHS} — avg loss: {avg_loss:.4f}")182 183 # Save checkpoint every epoch (remove mid-epoch checkpoint on success)184 ckpt_path = CHECKPOINT_DIR / f"checkpt_{epoch}.pt"185 torch.save({186 "epoch": epoch,187 "loss": avg_loss, # epoch avg loss — used by the plot cell188 "model": model.state_dict(),189 "optimiser": optimiser.state_dict(),190 "scaler": scaler.state_dict(),191 "best_loss": best_loss,192 "vocab_size": len(vocab),193 }, ckpt_path)194 mid = CHECKPOINT_DIR / "checkpt_mid.pt"195 if mid.exists():196 mid.unlink() # epoch complete — mid-epoch checkpoint no longer needed197 198 # Early stopping (matches original patience=2)199 if avg_loss < best_loss:200 best_loss = avg_loss201 patience_ctr = 0202 # Also save a named "best" checkpoint203 torch.save(torch.load(ckpt_path, weights_only=False), CHECKPOINT_DIR / "checkpt_best.pt")204 print(f" ✅ New best checkpoint saved.")205 else:206 patience_ctr += 1207 print(f" No improvement ({patience_ctr}/{PATIENCE})")208 if patience_ctr >= PATIENCE:209 print(f"Early stopping at epoch {epoch}.")210 break211 212 print(f"\nTraining complete. Checkpoints in: {CHECKPOINT_DIR}")213 214 215if __name__ == "__main__":216 parser = argparse.ArgumentParser()217 parser.add_argument("--resume", action="store_true",218 help="Resume from last checkpoint")219 args = parser.parse_args()220 train(resume=args.resume)221 