{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# 实验 08 · TinyGPT 端到端架构与自回归预训练 (Google Colab 交互式实验场)\n",
    "\n",
    "[![Open In Colab](https://colab.research.google.com/assets/colab-badge.svg)](https://colab.research.google.com/github/zenHeart/learn-llm/blob/main/notebooks/08_tinygpt_pretrain.ipynb)\n",
    "\n",
    "本 Notebook 专为配套 **Learn LLM** 第 8 章设计。利用你的 Google Pro 会员挂载高性能 GPU（A100 / L4 / T4），亲手跑通一个迷你 GPT 模型的端到端预训练与自回归生成循环！"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# [单元 1] 硬件加速器与显存检测\n",
    "import math\n",
    "import time\n",
    "import torch\n",
    "import torch.nn as nn\n",
    "import torch.nn.functional as F\n",
    "\n",
    "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n",
    "print(f\"运行环境: {device}\")\n",
    "if device == \"cuda\":\n",
    "    print(f\"🚀 算力加速卡: {torch.cuda.get_device_name(0)}\")\n",
    "    print(f\"可用显存: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB\")\n",
    "    torch.backends.cuda.matmul.allow_tf32 = True\n",
    "else:\n",
    "    print(\"⚠️ 建议切换至 GPU 运行时以大幅加速预训练过程！\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 架构实现：Transformer Block 与 TinyGPT"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "class TransformerBlock(nn.Module):\n",
    "    def __init__(self, d_model: int, num_heads: int, d_ff: int):\n",
    "        super().__init__()\n",
    "        self.ln1 = nn.LayerNorm(d_model)\n",
    "        self.attn = nn.MultiheadAttention(d_model, num_heads, batch_first=True)\n",
    "        self.ln2 = nn.LayerNorm(d_model)\n",
    "        self.ffn = nn.Sequential(\n",
    "            nn.Linear(d_model, d_ff),\n",
    "            nn.GELU(),\n",
    "            nn.Linear(d_ff, d_model)\n",
    "        )\n",
    "\n",
    "    def forward(self, x, is_causal: bool = True):\n",
    "        norm_x = self.ln1(x)\n",
    "        T = x.size(1)\n",
    "        mask = nn.Transformer.generate_square_subsequent_mask(T, device=x.device)\n",
    "        attn_out, _ = self.attn(norm_x, norm_x, norm_x, is_causal=True, attn_mask=mask)\n",
    "        x = x + attn_out\n",
    "        x = x + self.ffn(self.ln2(x))\n",
    "        return x\n",
    "\n",
    "class TinyGPT(nn.Module):\n",
    "    def __init__(self, vocab_size: int = 256, max_seq_len: int = 128, d_model: int = 64, num_heads: int = 4, num_layers: int = 2):\n",
    "        super().__init__()\n",
    "        self.max_seq_len = max_seq_len\n",
    "        self.tok_emb = nn.Embedding(vocab_size, d_model)\n",
    "        self.pos_emb = nn.Embedding(max_seq_len, d_model)\n",
    "        self.blocks = nn.ModuleList([\n",
    "            TransformerBlock(d_model, num_heads, d_ff=d_model * 4)\n",
    "            for _ in range(num_layers)\n",
    "        ])\n",
    "        self.ln_f = nn.LayerNorm(d_model)\n",
    "        self.head = nn.Linear(d_model, vocab_size, bias=False)\n",
    "\n",
    "    def forward(self, idx, targets=None):\n",
    "        B, T = idx.shape\n",
    "        pos = torch.arange(0, T, dtype=torch.long, device=idx.device)\n",
    "        x = self.tok_emb(idx) + self.pos_emb(pos)\n",
    "        for block in self.blocks:\n",
    "            x = block(x)\n",
    "        x = self.ln_f(x)\n",
    "        logits = self.head(x)\n",
    "        loss = None\n",
    "        if targets is not None:\n",
    "            loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1))\n",
    "        return logits, loss\n",
    "\n",
    "    @torch.no_grad()\n",
    "    def generate(self, idx, max_new_tokens: int = 20, temperature: float = 1.0):\n",
    "        for _ in range(max_new_tokens):\n",
    "            idx_cond = idx[:, -self.max_seq_len:]\n",
    "            logits, _ = self(idx_cond)\n",
    "            logits = logits[:, -1, :] / temperature\n",
    "            probs = F.softmax(logits, dim=-1)\n",
    "            next_tok = torch.multinomial(probs, num_samples=1)\n",
    "            idx = torch.cat((idx, next_tok), dim=1)\n",
    "        return idx"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 预训练循环与生成测试"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# 实例化模型并移至 GPU\n",
    "model = TinyGPT().to(device)\n",
    "optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)\n",
    "\n",
    "# 构造合成序列训练数据\n",
    "dummy_tokens = torch.randint(0, 256, (16, 32), device=device)\n",
    "x_batch = dummy_tokens[:, :-1]\n",
    "y_batch = dummy_tokens[:, 1:]\n",
    "\n",
    "print(\"开始执行预训练迭代...\")\n",
    "model.train()\n",
    "for step in range(10):\n",
    "    optimizer.zero_grad()\n",
    "    logits, loss = model(x_batch, y_batch)\n",
    "    loss.backward()\n",
    "    optimizer.step()\n",
    "    if (step + 1) % 2 == 0:\n",
    "        print(f\"Step {step+1:2d} | 交叉熵 Loss: {loss.item():.4f}\")\n",
    "\n",
    "# 执行自回归生成\n",
    "model.eval()\n",
    "prompt = torch.tensor([[65, 66, 67]], device=device) # 'ABC'\n",
    "generated = model.generate(prompt, max_new_tokens=10)\n",
    "print(f\"生成 Token 序列: {generated[0].tolist()}\")\n",
    "print(\"🎉 恭喜！TinyGPT 端到端训练与自回归生成循环验证通过！\")"
   ]
  }
 ],
 "metadata": {
  "accelerator": "GPU",
  "colab": {
   "provenance": []
  },
  "kernelspec": {
   "display_name": "Python 3",
   "name": "python3"
  },
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 0
}
