import re import math import random from collections import Counter from typing import Dict, List, Tuple import gradio as gr import numpy as np import torch from torch import nn from torch.utils.data import Dataset, DataLoader from datasets import load_dataset # ---------------------------- # Utilities # ---------------------------- SEED = 42 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) PAD, UNK = "", "" LABEL_NAMES = {0: "World", 1: "Sports", 2: "Business", 3: "Sci/Tech"} def simple_tokenize(text: str) -> List[str]: # lowercase + basic word tokenizer return re.findall(r"[A-Za-z0-9']+", text.lower()) def build_vocab(texts: List[str], max_vocab: int = 30000, min_freq: int = 2) -> Dict[str, int]: counter = Counter() for t in texts: counter.update(simple_tokenize(t)) # keep tokens by frequency up to max_vocab vocab = {PAD: 0, UNK: 1} for token, freq in counter.most_common(): if freq < min_freq: continue if len(vocab) >= max_vocab: break vocab[token] = len(vocab) return vocab def encode_text(text: str, vocab: Dict[str, int], max_len: int) -> List[int]: ids = [vocab.get(tok, vocab[UNK]) for tok in simple_tokenize(text)] if len(ids) >= max_len: return ids[:max_len] return ids + [vocab[PAD]] * (max_len - len(ids)) class AGNewsDataset(Dataset): def __init__(self, texts: List[str], labels: List[int], vocab: Dict[str, int], max_len: int): self.texts = texts self.labels = labels self.vocab = vocab self.max_len = max_len def __len__(self): return len(self.labels) def __getitem__(self, idx): x = torch.tensor(encode_text(self.texts[idx], self.vocab, self.max_len), dtype=torch.long) y = torch.tensor(self.labels[idx], dtype=torch.long) return x, y # ---------------------------- # Model # ---------------------------- class LSTMClassifier(nn.Module): def __init__(self, vocab_size: int, embed_dim: int, hidden_dim: int, num_layers: int, dropout: float, num_classes: int = 4, pad_idx: int = 0, bidirectional: bool = True): super().__init__() self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=pad_idx) self.lstm = nn.LSTM( input_size=embed_dim, hidden_size=hidden_dim, num_layers=num_layers, batch_first=True, bidirectional=bidirectional, dropout=dropout if num_layers > 1 else 0.0 ) out_dim = hidden_dim * (2 if bidirectional else 1) self.dropout = nn.Dropout(dropout) self.fc = nn.Linear(out_dim, num_classes) def forward(self, x): emb = self.embedding(x) # (B, T, E) outputs, (h_n, _) = self.lstm(emb) # h_n: (num_layers*num_dir, B, H) # Take last layer's final hidden state(s) if self.lstm.bidirectional: h_last = torch.cat([h_n[-2], h_n[-1]], dim=1) # (B, 2H) else: h_last = h_n[-1] # (B, H) logits = self.fc(self.dropout(h_last)) return logits # ---------------------------- # Training / Eval helpers # ---------------------------- def accuracy_from_logits(logits: torch.Tensor, y: torch.Tensor) -> float: preds = logits.argmax(dim=1) return (preds == y).float().mean().item() def train_one_epoch(model, loader, optimizer, criterion, device): model.train() total_loss, total_acc, n = 0.0, 0.0, 0 for x, y in loader: x, y = x.to(device), y.to(device) optimizer.zero_grad() logits = model(x) loss = criterion(logits, y) loss.backward() optimizer.step() bsz = y.size(0) total_loss += loss.item() * bsz total_acc += accuracy_from_logits(logits, y) * bsz n += bsz return total_loss / n, total_acc / n def evaluate(model, loader, criterion, device): model.eval() total_loss, total_acc, n = 0.0, 0.0, 0 with torch.no_grad(): for x, y in loader: x, y = x.to(device), y.to(device) logits = model(x) loss = criterion(logits, y) bsz = y.size(0) total_loss += loss.item() * bsz total_acc += accuracy_from_logits(logits, y) * bsz n += bsz return total_loss / n, total_acc / n # ---------------------------- # Gradio state container # ---------------------------- class AppState: def __init__(self): self.vocab = None self.max_len = None self.model = None self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") self.train_loader = None self.valid_loader = None self.test_loader = None # ---------------------------- # Gradio actions # ---------------------------- def setup_data(max_vocab: int, max_len: int, batch_size: int, use_full: bool): """Load AG News, build vocab, and create DataLoaders.""" state = AppState() ds = load_dataset("ag_news") # (Optional) reduce size for quicker demos if not use_full: # ~24k train rows -> take 12k for faster runs ds["train"] = ds["train"].shuffle(seed=SEED).select(range(12000)) ds["test"] = ds["test"].shuffle(seed=SEED).select(range(2400)) # build vocab on train texts train_texts = [r["text"] for r in ds["train"]] train_labels = [int(r["label"]) for r in ds["train"]] vocab = build_vocab(train_texts, max_vocab=max_vocab, min_freq=2) # split train into train/valid (90/10) split = int(0.9 * len(train_texts)) tr_texts, va_texts = train_texts[:split], train_texts[split:] tr_labels, va_labels = train_labels[:split], train_labels[split:] state.vocab = vocab state.max_len = max_len train_ds = AGNewsDataset(tr_texts, tr_labels, vocab, max_len) valid_ds = AGNewsDataset(va_texts, va_labels, vocab, max_len) test_ds = AGNewsDataset([r["text"] for r in ds["test"]], [int(r["label"]) for r in ds["test"]], vocab, max_len) state.train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True) state.valid_loader = DataLoader(valid_ds, batch_size=batch_size) state.test_loader = DataLoader(test_ds, batch_size=batch_size) vocab_size = len(vocab) msg = ( f"✅ Data ready\n" f"- Train: {len(train_ds)} | Valid: {len(valid_ds)} | Test: {len(test_ds)}\n" f"- Vocab size: {vocab_size}\n" f"- Max length: {max_len}\n" f"- Batch size: {batch_size}\n" f"- Using {'FULL' if use_full else 'REDUCED'} dataset" ) return state, msg def train(state: AppState, embed_dim: int, hidden_dim: int, num_layers: int, dropout: float, epochs: int, lr: float): if state is None or state.vocab is None: return None, "❗Please run Setup first." model = LSTMClassifier( vocab_size=len(state.vocab), embed_dim=embed_dim, hidden_dim=hidden_dim, num_layers=num_layers, dropout=dropout, num_classes=4, pad_idx=0, bidirectional=True ).to(state.device) opt = torch.optim.Adam(model.parameters(), lr=lr) criterion = nn.CrossEntropyLoss() history_rows = [["epoch", "train_loss", "train_acc", "valid_loss", "valid_acc"]] best_val = -1 for ep in range(1, epochs + 1): tr_loss, tr_acc = train_one_epoch(model, state.train_loader, opt, criterion, state.device) va_loss, va_acc = evaluate(model, state.valid_loader, criterion, state.device) history_rows.append([ep, round(tr_loss, 4), round(tr_acc, 4), round(va_loss, 4), round(va_acc, 4)]) if va_acc > best_val: best_val = va_acc state.model = model table_md = "| " + " | ".join(history_rows[0]) + " |\n|---|---:|---:|---:|---:|\n" for r in history_rows[1:]: table_md += "| " + " | ".join(str(x) for x in r) + " |\n" table_md += f"\n**Best validation accuracy:** {best_val:.4f}" return state, table_md def test_eval(state: AppState): if state is None or state.model is None: return "❗Please train the model first." criterion = nn.CrossEntropyLoss() te_loss, te_acc = evaluate(state.model, state.test_loader, criterion, state.device) return f"🧪 Test Loss: {te_loss:.4f} | Test Accuracy: {te_acc:.4f}" def predict(state: AppState, text: str): if state is None or state.model is None or state.vocab is None: return "❗Please train the model first." state.model.eval() with torch.no_grad(): x = torch.tensor([encode_text(text, state.vocab, state.max_len)], dtype=torch.long).to(state.device) logits = state.model(x) probs = torch.softmax(logits, dim=1).cpu().numpy().flatten() top = int(probs.argmax()) label = LABEL_NAMES[top] conf = float(probs[top]) dist = {LABEL_NAMES[i]: float(p) for i, p in enumerate(probs)} return f"Prediction: **{label}** (confidence {conf:.3f})", dist # ---------------------------- # Gradio UI # ---------------------------- with gr.Blocks(title="AG News LSTM (PyTorch + 🤗 Datasets)") as demo: gr.Markdown("# AG News Topic Classification (LSTM)\nLoad & preprocess with 🤗 Datasets → Train an LSTM in PyTorch → Evaluate → Predict") state = gr.State() with gr.Tab("1) Setup & Preprocess"): with gr.Row(): max_vocab = gr.Slider(5_000, 80_000, value=30_000, step=1_000, label="Max Vocab Size") max_len = gr.Slider(32, 512, value=128, step=8, label="Max Sequence Length") batch_sz = gr.Slider(8, 128, value=64, step=8, label="Batch Size") use_full = gr.Checkbox(False, label="Use FULL dataset (unchecked uses a reduced subset for speed)") setup_btn = gr.Button("Setup Data") setup_out = gr.Markdown() setup_btn.click( setup_data, inputs=[max_vocab, max_len, batch_sz, use_full], outputs=[state, setup_out] ) with gr.Tab("2) Train"): with gr.Row(): embed_dim = gr.Slider(32, 512, value=128, step=16, label="Embedding Dim") hidden_dim = gr.Slider(32, 512, value=128, step=16, label="Hidden Dim") num_layers = gr.Slider(1, 4, value=1, step=1, label="LSTM Layers") with gr.Row(): dropout = gr.Slider(0.0, 0.6, value=0.3, step=0.05, label="Dropout") epochs = gr.Slider(1, 15, value=3, step=1, label="Epochs") lr = gr.Number(value=2e-3, label="Learning Rate") train_btn = gr.Button("Train") history_md = gr.Markdown() train_btn.click( train, inputs=[state, embed_dim, hidden_dim, num_layers, dropout, epochs, lr], outputs=[state, history_md] ) with gr.Tab("3) Evaluate"): test_btn = gr.Button("Evaluate on Test Split") test_out = gr.Markdown() test_btn.click(test_eval, inputs=[state], outputs=[test_out]) with gr.Tab("4) Predict"): user_txt = gr.Textbox(lines=4, label="Paste a news headline or short article") pred_btn = gr.Button("Classify") pred_label = gr.Markdown() pred_probs = gr.Label(num_top_classes=4) pred_btn.click(predict, inputs=[state, user_txt], outputs=[pred_label, pred_probs]) if __name__ == "__main__": demo.launch()