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