{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# 实验 03 · 从零实现标量自动微分引擎 (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/03_scalar_autograd.ipynb)\n",
    "\n",
    "本 Notebook 专为配套 **Learn LLM** 第 3 章设计。读者可以在此手写微型标量 autograd 引擎 `Value`，并通过拓扑排序实现自动反向传播。"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# [单元 1] 运行环境准备\n",
    "import math\n",
    "import torch\n",
    "\n",
    "print(f\"PyTorch 版本: {torch.__version__}\")\n",
    "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n",
    "print(f\"当前硬件设备: {device}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 核心实现：标量节点 `Value` 与闭包 `_backward`"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "class Value:\n",
    "    \"\"\"标量微积分计算图节点\"\"\"\n",
    "    def __init__(self, data: float, _children: tuple = (), _op: str = ''):\n",
    "        self.data = float(data)\n",
    "        self.grad = 0.0\n",
    "        self._backward = lambda: None\n",
    "        self._prev = set(_children)\n",
    "        self._op = _op\n",
    "\n",
    "    def __repr__(self):\n",
    "        return f\"Value(data={self.data:.4f}, grad={self.grad:.4f})\"\n",
    "\n",
    "    def __add__(self, other):\n",
    "        other = other if isinstance(other, Value) else Value(other)\n",
    "        out = Value(self.data + other.data, (self, other), '+')\n",
    "        def _backward():\n",
    "            self.grad += 1.0 * out.grad\n",
    "            other.grad += 1.0 * out.grad\n",
    "        out._backward = _backward\n",
    "        return out\n",
    "\n",
    "    def __mul__(self, other):\n",
    "        other = other if isinstance(other, Value) else Value(other)\n",
    "        out = Value(self.data * other.data, (self, other), '*')\n",
    "        def _backward():\n",
    "            self.grad += other.data * out.grad\n",
    "            other.grad += self.data * out.grad\n",
    "        out._backward = _backward\n",
    "        return out\n",
    "\n",
    "    def tanh(self):\n",
    "        t = math.tanh(self.data)\n",
    "        out = Value(t, (self,), 'tanh')\n",
    "        def _backward():\n",
    "            self.grad += (1.0 - t**2) * out.grad\n",
    "        out._backward = _backward\n",
    "        return out\n",
    "\n",
    "    def backward(self):\n",
    "        # 逆拓扑排序遍历计算图\n",
    "        topo = []\n",
    "        visited = set()\n",
    "        def build_topo(v):\n",
    "            if v not in visited:\n",
    "                visited.add(v)\n",
    "                for child in v._prev:\n",
    "                    build_topo(child)\n",
    "                topo.append(v)\n",
    "        build_topo(self)\n",
    "        self.grad = 1.0\n",
    "        for node in reversed(topo):\n",
    "            node._backward()"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 验证对拍：与 PyTorch 原生 Autograd 严格对齐"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# 1. 手写引擎前向与反向\n",
    "x1 = Value(2.0)\n",
    "x2 = Value(0.0)\n",
    "w1 = Value(-3.0)\n",
    "w2 = Value(1.0)\n",
    "b = Value(6.8813735870195432)\n",
    "n = x1 * w1 + x2 * w2 + b\n",
    "o = n.tanh()\n",
    "o.backward()\n",
    "\n",
    "# 2. PyTorch 原生计算\n",
    "tx1 = torch.tensor(2.0, requires_grad=True, dtype=torch.float64)\n",
    "tx2 = torch.tensor(0.0, requires_grad=True, dtype=torch.float64)\n",
    "tw1 = torch.tensor(-3.0, requires_grad=True, dtype=torch.float64)\n",
    "tw2 = torch.tensor(1.0, requires_grad=True, dtype=torch.float64)\n",
    "tb = torch.tensor(6.8813735870195432, requires_grad=True, dtype=torch.float64)\n",
    "tn = tx1 * tw1 + tx2 * tw2 + tb\n",
    "to = torch.tanh(tn)\n",
    "to.backward()\n",
    "\n",
    "print(f\"输出值对比: Value={o.data:.8f}, Torch={to.item():.8f}\")\n",
    "print(f\"x1.grad 对比: Value={x1.grad:.8f}, Torch={tx1.grad.item():.8f}\")\n",
    "print(f\"w1.grad 对比: Value={w1.grad:.8f}, Torch={tw1.grad.item():.8f}\")\n",
    "\n",
    "assert math.isclose(o.data, to.item(), rel_tol=1e-6)\n",
    "assert math.isclose(x1.grad, tx1.grad.item(), rel_tol=1e-6)\n",
    "assert math.isclose(w1.grad, tw1.grad.item(), rel_tol=1e-6)\n",
    "print(\"🎉 恭喜！手写标量 autograd 与 PyTorch 精度 100% 对齐！\")"
   ]
  }
 ],
 "metadata": {
  "accelerator": "GPU",
  "colab": {
   "provenance": []
  },
  "kernelspec": {
   "display_name": "Python 3",
   "name": "python3"
  },
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 0
}
