{
 "cells": [
  {
   "cell_type": "markdown",
   "id": "bigram-goal",
   "metadata": {},
   "source": [
    "# Bigram language model\n",
    "\n",
    "Prepare the Tiny Shakespeare token stream. The next step is to build context-window `(x, y)` samples."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 1,
   "id": "setup",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-08-23T00:04:35.658566Z",
     "iopub.status.busy": "2026-08-23T00:04:35.658362Z",
     "iopub.status.idle": "2026-08-23T00:04:35.666160Z",
     "shell.execute_reply": "2026-08-23T00:04:35.665287Z"
    }
   },
   "outputs": [],
   "source": [
    "import sys\n",
    "from pathlib import Path\n",
    "\n",
    "repo_root = Path.cwd()\n",
    "if not (repo_root / \"data\").exists():\n",
    "    repo_root = repo_root.parent\n",
    "sys.path.insert(0, str(repo_root))\n",
    "\n",
    "from src.dataset import get_batch, load_tiny_shakespeare_tokens, split_token_stream\n",
    "from src.tokenizer import bpe_decode"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "load-data",
   "metadata": {},
   "source": [
    "## Load and encode Tiny Shakespeare"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 2,
   "id": "encode-data",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-08-23T00:04:35.668553Z",
     "iopub.status.busy": "2026-08-23T00:04:35.668380Z",
     "iopub.status.idle": "2026-08-23T00:04:52.013605Z",
     "shell.execute_reply": "2026-08-23T00:04:52.009499Z"
    }
   },
   "outputs": [],
   "source": [
    "tokens, vocab, _ = load_tiny_shakespeare_tokens(repo_root / \"data\")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "split-data",
   "metadata": {},
   "source": [
    "## Split the token stream"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 3,
   "id": "split",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-08-23T00:04:52.019618Z",
     "iopub.status.busy": "2026-08-23T00:04:52.019448Z",
     "iopub.status.idle": "2026-08-23T00:04:52.025319Z",
     "shell.execute_reply": "2026-08-23T00:04:52.025019Z"
    }
   },
   "outputs": [],
   "source": [
    "train_tokens, validation_tokens = split_token_stream(tokens)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "check-data",
   "metadata": {},
   "source": [
    "## Sanity check"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 4,
   "id": "sanity-check",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-08-23T00:04:52.026468Z",
     "iopub.status.busy": "2026-08-23T00:04:52.026407Z",
     "iopub.status.idle": "2026-08-23T00:04:52.028727Z",
     "shell.execute_reply": "2026-08-23T00:04:52.028383Z"
    }
   },
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Total tokens: 590,498\n",
      "Training tokens: 531,448\n",
      "Validation tokens: 59,050\n"
     ]
    }
   ],
   "source": [
    "print(f\"Total tokens: {len(tokens):,}\")\n",
    "print(f\"Training tokens: {len(train_tokens):,}\")\n",
    "print(f\"Validation tokens: {len(validation_tokens):,}\")\n",
    "# print(bpe_decode(train_tokens[:100], vocab))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 5,
   "id": "8590a86a",
   "metadata": {},
   "outputs": [],
   "source": [
    "import random\n",
    "\n",
    "# Set parameters for batching\n",
    "seed = 42\n",
    "random.seed(seed)\n",
    "block_size = 8\n",
    "batch_size = 32\n",
    "\n",
    "\n",
    "def sample_batch(split):\n",
    "    return get_batch(\n",
    "        split,\n",
    "        train_tokens,\n",
    "        validation_tokens,\n",
    "        block_size,\n",
    "        batch_size,\n",
    "    )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 6,
   "id": "d9eeeccf",
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "from torch import nn\n",
    "\n",
    "vocab_size = len(vocab)\n",
    "n_embd = 32\n",
    "\n",
    "\n",
    "class BigramLanguageModel(nn.Module):\n",
    "    def __init__(self, vocab_size, n_embd):\n",
    "        super().__init__()\n",
    "        self.token_embedding_table = nn.Embedding(vocab_size, n_embd)\n",
    "        self.lm_head = nn.Linear(n_embd, vocab_size)\n",
    "\n",
    "    def forward(self, idx, targets=None):\n",
    "        x = self.token_embedding_table(idx)\n",
    "        logits = self.lm_head(x)\n",
    "\n",
    "        if targets is None:\n",
    "            loss = None\n",
    "        else:\n",
    "            B, T, V = logits.shape\n",
    "            logits = logits.view(B * T, V)\n",
    "            targets = targets.view(B * T)\n",
    "            loss = nn.functional.cross_entropy(logits, targets)\n",
    "\n",
    "        return logits, loss\n",
    "\n",
    "    def generate(self, idx, max_new_tokens):\n",
    "        for _ in range(max_new_tokens):\n",
    "            logits, _ = self(idx)\n",
    "            logits = logits[:, -1, :]\n",
    "            probs = torch.softmax(logits, dim=-1)\n",
    "            idx_next = torch.multinomial(probs, num_samples=1)\n",
    "            idx = torch.cat((idx, idx_next), dim=1)\n",
    "\n",
    "        return idx"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 7,
   "id": "5670bee5",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Step 0: loss = 6.4571\n",
      "Step 100: loss = 6.1014\n",
      "Step 200: loss = 5.7985\n",
      "Step 300: loss = 5.5250\n",
      "Step 400: loss = 5.2603\n",
      "Step 500: loss = 5.1149\n",
      "Step 600: loss = 4.8774\n",
      "Step 700: loss = 4.6288\n",
      "Step 800: loss = 4.5432\n",
      "Step 900: loss = 4.4973\n"
     ]
    }
   ],
   "source": [
    "model = BigramLanguageModel(vocab_size, n_embd)\n",
    "optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)\n",
    "\n",
    "for step in range(1000):\n",
    "    x_batch, y_batch = sample_batch(\"train\")\n",
    "    x_batch = torch.tensor(x_batch, dtype=torch.long)\n",
    "    y_batch = torch.tensor(y_batch, dtype=torch.long)\n",
    "\n",
    "    logits, loss = model(x_batch, y_batch)\n",
    "\n",
    "    optimizer.zero_grad(set_to_none=True)\n",
    "    loss.backward()\n",
    "    optimizer.step()\n",
    "\n",
    "    if step % 100 == 0:\n",
    "        print(f\"Step {step}: loss = {loss.item():.4f}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 8,
   "id": "4766adfe",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "validation loss: 4.3758955001831055\n",
      "log vocab size: 6.238324625039508\n"
     ]
    },
    {
     "data": {
      "text/plain": [
       "BigramLanguageModel(\n",
       "  (token_embedding_table): Embedding(512, 32)\n",
       "  (lm_head): Linear(in_features=32, out_features=512, bias=True)\n",
       ")"
      ]
     },
     "execution_count": 8,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "import math\n",
    "\n",
    "model.eval()\n",
    "\n",
    "with torch.no_grad():\n",
    "    xb, yb = sample_batch(\"validation\")\n",
    "\n",
    "    xb = torch.tensor(xb, dtype=torch.long)\n",
    "    yb = torch.tensor(yb, dtype=torch.long)\n",
    "\n",
    "    _, val_loss = model(xb, yb)\n",
    "\n",
    "print(\"validation loss:\", val_loss.item())\n",
    "print(\"log vocab size:\", math.log(len(vocab)))\n",
    "\n",
    "model.train()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 9,
   "id": "cb719c2d",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "\u0000�d\n",
      "ookPAE:\n",
      "Oerhour,\n",
      "R:\n",
      "Tll on , and the e to to youPant shs ow, behoughghton neygmeweie emcyour aanour  and such'd I git my atOr ce, syspbuINGAORove:\n",
      "ut aupnoRLence:\n",
      "NG that sece �y bebla er \n"
     ]
    }
   ],
   "source": [
    "context = torch.tensor([[0]], dtype=torch.long)\n",
    "out = model.generate(context, max_new_tokens=100)\n",
    "print(bpe_decode(out[0].tolist(), vocab, errors=\"replace\"))"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "lmlab (3.12.13.final.0)",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.12.13"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
