{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# 实验 10 · 手写 LoRA 低秩微调与零时延权重合并 (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/10_sft_lora.ipynb)\n",
    "\n",
    "本 Notebook 专为配套 **Learn LLM** 第 10 章设计。在 Google Colab Pro 上手写 LoRA 核心机制，理解低秩约束、参数冻结以及零延迟权重合并（Weight Merge）的工业级原理！"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# [单元 1] 硬件加速器与显存检测\n",
    "import math\n",
    "import torch\n",
    "import torch.nn as nn\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"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 核心机制：LoRALinear 适配器层\n",
    "\n",
    "$$h = W_0 x + \\frac{\\alpha}{r} B A x$$\n",
    "\n",
    "- $W_0 \\in \\mathbb{R}^{d_{out} \\times d_{in}}$ 彻底冻结（`requires_grad=False`）；\n",
    "- $A \\in \\mathbb{R}^{r \\times d_{in}}$ 高斯初始化；\n",
    "- $B \\in \\mathbb{R}^{d_{out} \\times r}$ 全零初始化，**保证训练初始步 $\\Delta W = 0$，严格维持原模型基线输出！**"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "class LoRALinear(nn.Module):\n",
    "    def __init__(self, in_features: int, out_features: int, r: int = 8, lora_alpha: float = 16.0):\n",
    "        super().__init__()\n",
    "        self.in_features = in_features\n",
    "        self.out_features = out_features\n",
    "        self.r = r\n",
    "        self.scaling = lora_alpha / r\n",
    "        \n",
    "        # 原始主干权重 (冻结)\n",
    "        self.weight = nn.Parameter(torch.randn(out_features, in_features), requires_grad=False)\n",
    "        \n",
    "        # 低秩适配矩阵 (可训参数量仅约 1%)\n",
    "        self.lora_A = nn.Parameter(torch.randn(r, in_features) / math.sqrt(r))\n",
    "        self.lora_B = nn.Parameter(torch.zeros(out_features, r)) # 零初始化至关重要\n",
    "        \n",
    "    def forward(self, x):\n",
    "        # 基础前向\n",
    "        base_out = nn.functional.linear(x, self.weight)\n",
    "        # 旁路低秩前向\n",
    "        lora_out = (x @ self.lora_A.T) @ self.lora_B.T * self.scaling\n",
    "        return base_out + lora_out\n",
    "        \n",
    "    def merge_weights(self):\n",
    "        \"\"\"部署推理时的零时延权重合并 (Zero-latency Merge)\"\"\"\n",
    "        delta_w = (self.lora_B @ self.lora_A) * self.scaling\n",
    "        self.weight.data += delta_w.data\n",
    "        # 合并后即可丢弃 A 和 B，恢复为单一稠密矩阵！\n",
    "        print(\"✅ 权重合并成功，推理时无任何额外计算开销！\")\n",
    "\n",
    "# 验证零初始偏移特性\n",
    "lora_layer = LoRALinear(128, 64, r=4).to(device)\n",
    "x = torch.randn(2, 128, device=device)\n",
    "\n",
    "base_only = nn.functional.linear(x, lora_layer.weight)\n",
    "lora_init = lora_layer(x)\n",
    "\n",
    "diff = torch.norm(lora_init - base_only).item()\n",
    "print(f\"初始输出偏差: {diff:.8f}\")\n",
    "assert diff == 0.0, \"初始步 LoRA 输出必须严格等于基础权重输出！\"\n",
    "print(\"✅ 零初始化基线保持特性验证通过！\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 零时延部署验证：合并权重后前后计算一致性"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# 模拟微调：给 B 赋予微小更新\n",
    "lora_layer.lora_B.data.add_(0.01)\n",
    "\n",
    "out_before_merge = lora_layer(x)\n",
    "lora_layer.merge_weights()\n",
    "out_after_merge = nn.functional.linear(x, lora_layer.weight)\n",
    "\n",
    "merge_diff = torch.norm(out_before_merge - out_after_merge).item()\n",
    "print(f\"合并前后数值差异: {merge_diff:.8f}\")\n",
    "assert math.isclose(merge_diff, 0.0, abs_tol=1e-5)\n",
    "print(\"🎉 恭喜！LoRA 低秩微调与零时延权重合并检验全部 PASS！\")"
   ]
  }
 ],
 "metadata": {
  "accelerator": "GPU",
  "colab": {
   "provenance": []
  },
  "kernelspec": {
   "display_name": "Python 3",
   "name": "python3"
  },
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 0
}
