{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# 实验 05：从零实现 LoRA\n",
    "\n",
    "**运行条件：** CPU，约 1 分钟，内存低于 1 GB。\n",
    "\n",
    "LoRA 冻结原始权重，只训练低秩矩阵 $A$ 和 $B$。本实验在合成任务上验证参数量、训练过程和权重合并。"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "import torch.nn as nn\n",
    "import matplotlib.pyplot as plt\n",
    "\n",
    "torch.manual_seed(23)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "class LoRALinear(nn.Module):\n",
    "    def __init__(self, base: nn.Linear, rank=4, alpha=8):\n",
    "        super().__init__()\n",
    "        self.base = base\n",
    "        self.base.requires_grad_(False)\n",
    "        self.A = nn.Parameter(torch.empty(rank, base.in_features))\n",
    "        self.B = nn.Parameter(torch.zeros(base.out_features, rank))\n",
    "        nn.init.kaiming_uniform_(self.A, a=5**0.5)\n",
    "        self.scale = alpha / rank\n",
    "    def forward(self, x):\n",
    "        return self.base(x) + (x @ self.A.T @ self.B.T) * self.scale\n",
    "    def merged_weight(self):\n",
    "        return self.base.weight + (self.B @ self.A) * self.scale\n",
    "\n",
    "base = nn.Linear(64, 32, bias=False)\n",
    "layer = LoRALinear(base, rank=4, alpha=8)\n",
    "trainable = sum(p.numel() for p in layer.parameters() if p.requires_grad)\n",
    "full = base.weight.numel()\n",
    "print('full parameters:', full)\n",
    "print('LoRA trainable:', trainable, f'({100*trainable/full:.1f}%)')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 训练一个低秩增量\n",
    "\n",
    "目标映射由冻结基座权重加一个 rank-4 增量构造，因此 LoRA 可以在足够训练后恢复它。"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "true_A = torch.randn(4, 64) * 0.12\n",
    "true_B = torch.randn(32, 4) * 0.12\n",
    "features = torch.randn(512, 64)\n",
    "with torch.no_grad():\n",
    "    targets = base(features) + features @ true_A.T @ true_B.T\n",
    "optimizer = torch.optim.AdamW([layer.A, layer.B], lr=0.03)\n",
    "losses = []\n",
    "for step in range(240):\n",
    "    loss = ((layer(features) - targets) ** 2).mean()\n",
    "    optimizer.zero_grad(set_to_none=True); loss.backward(); optimizer.step()\n",
    "    losses.append(loss.item())\n",
    "print('initial loss:', losses[0], 'final loss:', losses[-1])\n",
    "plt.plot(losses); plt.yscale('log'); plt.xlabel('step'); plt.ylabel('MSE'); plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 合并权重\n",
    "\n",
    "部署时可以把低秩增量合并进基座权重，消除额外矩阵乘法。"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "merged = nn.Linear(64, 32, bias=False)\n",
    "with torch.no_grad(): merged.weight.copy_(layer.merged_weight())\n",
    "test = torch.randn(8, 64)\n",
    "error = (layer(test) - merged(test)).abs().max().item()\n",
    "print('merge error:', error)\n",
    "assert error < 1e-5"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## 练习\n",
    "\n",
    "1. 把目标增量的 rank 改成 16，而 LoRA rank 保持 4，观察误差下限。\n",
    "2. 比较 rank 为 1、2、4、8 时的参数量和收敛速度。\n",
    "3. 解释为什么 `B` 初始化为零可以确保训练开始时模型行为不变。\n",
    "4. 将 LoRA 放到 Attention 的 Q、K、V 和输出投影，比较可训练参数量。"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "name": "python",
   "version": "3.11"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
