CoolFace
Apppublic

JetLaggedByData/scifi-forge

sourceHugging Facemitupdated 5mo agoView on Hugging Face
0likes
train.py221 linesDownload Raw Back to v1_baseline
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