{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# 实验 07 · 多头因果自注意力机制 (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/07_attention_mechanism.ipynb)\n",
    "\n",
    "本 Notebook 专为配套 **Learn LLM** 第 7 章设计。利用你的 Google Pro 会员挂载高性能 GPU（A100 / L4 / T4），亲手实现 Transformer 核心的 Scaled Dot-Product Attention 与 Multi-Head Attention！"
   ]
  },
  {
   "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\"🚀 GPU 型号: {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(\"⚠️ 建议点击顶部「修改运行时类型」，选择 T4 或 A100 GPU 开启全速模式！\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 核心算法：缩放点积因果自注意力 (Scaled Dot-Product Attention)\n",
    "\n",
    "$$\\mathrm{Attention}(Q, K, V) = \\mathrm{softmax}\\left(\\frac{QK^T}{\\sqrt{d_k}} + M\\right)V$$"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def scaled_dot_product_attention(Q, K, V, is_causal: bool = True):\n",
    "    \"\"\"\n",
    "    Q: (B, H, T, d_k)\n",
    "    K: (B, H, T, d_k)\n",
    "    V: (B, H, T, d_v)\n",
    "    \"\"\"\n",
    "    d_k = Q.size(-1)\n",
    "    # 1. 点积并除以缩放因子 sqrt(d_k)\n",
    "    scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)\n",
    "    \n",
    "    # 2. 因果下三角遮罩 (防止未来信息泄露)\n",
    "    if is_causal:\n",
    "        T = Q.size(-2)\n",
    "        mask = torch.triu(torch.full((T, T), float('-inf'), device=Q.device), diagonal=1)\n",
    "        scores = scores + mask\n",
    "        \n",
    "    # 3. Softmax 概率归一化\n",
    "    weights = F.softmax(scores, dim=-1)\n",
    "    \n",
    "    # 4. 汇聚 Value\n",
    "    out = torch.matmul(weights, V)\n",
    "    return out, weights\n",
    "\n",
    "# 验证形状与遮罩行为\n",
    "B, H, T, D = 2, 4, 8, 32\n",
    "q = torch.randn(B, H, T, D, device=device)\n",
    "k = torch.randn(B, H, T, D, device=device)\n",
    "v = torch.randn(B, H, T, D, device=device)\n",
    "\n",
    "out, weights = scaled_dot_product_attention(q, k, v, is_causal=True)\n",
    "print(f\"输出张量形状: {out.shape} (预期 [{B}, {H}, {T}, {D}])\")\n",
    "assert out.shape == (B, H, T, D)\n",
    "# 检查因果遮罩严格性：右上角权重必须精确为 0.0\n",
    "assert weights[0, 0, 0, 1].item() == 0.0\n",
    "print(\"✅ 缩放点积因果自注意力测试通过！\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 完整模块：多头自注意力层 (CausalMultiHeadAttention)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "class CausalMultiHeadAttention(nn.Module):\n",
    "    def __init__(self, d_model: int, num_heads: int):\n",
    "        super().__init__()\n",
    "        assert d_model % num_heads == 0, \"d_model 必须能被 num_heads 整除\"\n",
    "        self.d_model = d_model\n",
    "        self.num_heads = num_heads\n",
    "        self.d_k = d_model // num_heads\n",
    "        \n",
    "        # 线性投射层\n",
    "        self.q_proj = nn.Linear(d_model, d_model)\n",
    "        self.k_proj = nn.Linear(d_model, d_model)\n",
    "        self.v_proj = nn.Linear(d_model, d_model)\n",
    "        self.out_proj = nn.Linear(d_model, d_model)\n",
    "\n",
    "    def forward(self, x):\n",
    "        B, T, C = x.shape\n",
    "        # 投射并拆分头 (B, T, C) -> (B, H, T, d_k)\n",
    "        q = self.q_proj(x).view(B, T, self.num_heads, self.d_k).transpose(1, 2)\n",
    "        k = self.k_proj(x).view(B, T, self.num_heads, self.d_k).transpose(1, 2)\n",
    "        v = self.v_proj(x).view(B, T, self.num_heads, self.d_k).transpose(1, 2)\n",
    "        \n",
    "        out, _ = scaled_dot_product_attention(q, k, v, is_causal=True)\n",
    "        \n",
    "        # 拼接多头并线性变换回 (B, T, C)\n",
    "        out = out.transpose(1, 2).contiguous().view(B, T, C)\n",
    "        return self.out_proj(out)\n",
    "\n",
    "# 实例化并在 GPU 上测试吞吐\n",
    "mha = CausalMultiHeadAttention(d_model=128, num_heads=4).to(device)\n",
    "dummy_input = torch.randn(8, 64, 128, device=device)\n",
    "output = mha(dummy_input)\n",
    "print(f\"MHA 前向输出形状: {output.shape}\")\n",
    "assert output.shape == dummy_input.shape\n",
    "print(\"🎉 恭喜！完整多头自注意力层测试全部 PASS！\")"
   ]
  }
 ],
 "metadata": {
  "accelerator": "GPU",
  "colab": {
   "provenance": []
  },
  "kernelspec": {
   "display_name": "Python 3",
   "name": "python3"
  },
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 0
}
