第 13 章 2023—2024:高效大模型——LoRA、MoE 与 Mamba
学习目标
- 理解 LoRA(2021/2023 流行)的低秩适配思想,用少量参数微调大模型
- 理解混合专家(MoE)的稀疏激活与负载均衡
- 理解 Mamba(2023)的选择性状态空间模型
- 分别用 PyTorch 实现三种架构
13.1 时代背景:大模型很贵,怎么省钱
2023—2024 年,大模型从「能不能做」变成「能不能用得起」:1750 亿参数的 GPT-3 微调一轮,权重更新就是数百 GB;推理时每个 token 都要过全部参数。三个方向应运而生:
- LoRA:微调时只训练小矩阵(参数减少 100 倍以上);
- MoE:推理时只激活部分专家(计算量不变,容量扩大);
- Mamba:用线性递归替代注意力(长序列推理省内存)。
13.2 LoRA:训练小矩阵,冻结大模型
低秩适配(Low-Rank Adaptation, LoRA)的核心假设:微调时,权重变化 ΔW 是低秩的。于是不直接更新 W,而是把它写成两个小矩阵的乘积:
W' = W + ΔW = W + (B·Aᵀ)·scaleA 是 (in, r),B 是 (out, r),秩 r 通常取 4~64。训练时 W 冻结,只更新 A、B。一个 1750 亿参数模型,LoRA 只训练几百万参数。
examples/ch13_lora.py 在 MNIST 奇偶分类上对比全量微调与 LoRA:
全量微调:可训练参数 50,370,测试准确率 93.7%
LoRA 微调:可训练参数 3,522,测试准确率 87.8%
LoRA 增量矩阵 ||ΔW||/||W|| = 0.4058LoRA 只训练 7% 的参数,准确率达到全量微调的 94%。||ΔW||/||W|| = 0.41 说明权重确实只发生了小幅低秩修正。真实世界里,LoRA 让「一张消费级显卡微调 70B 模型」成为可能,也催生了成千上万个社区微调模型。
13.3 MoE:让专家各司其职
混合专家(Mixture of Experts, MoE)的思想:与其让一个巨大的 FFN 处理所有输入,不如准备 N 个专家网络,由路由器决定每个输入交给哪几个专家(top-2)。
好处:模型总参数量巨大(容量大),但每个 token 只激活 2 个专家(计算量可控)。Switch Transformer(2021)与 Mixtral(2024)把这条路带到主流。
MoE 的关键工程问题是负载均衡:路由器可能「偏爱」少数专家。解决方法是加辅助损失,惩罚不均衡的路由分布。
examples/ch13_moe.py 用 4 个专家 + top-2 路由做二分类:
MoE 参数量:17,062
step 100:loss = 0.5332,准确率 = 78.1%
step 200:loss = 0.5633,准确率 = 80.5%
step 300:loss = 0.4904,准确率 = 82.8%
step 400:loss = 0.5492,准确率 = 79.7%
step 500:loss = 0.5713,准确率 = 78.9%
各专家被激活次数: [33934, 35415, 31661, 26990] → 负载均衡损失让使用接近均匀4 个专家的使用次数 3.2 万~3.5 万,相当均匀——负载均衡损失生效了。对比无辅助损失时(某个专家可能被用 2 倍),可以看到 MoE 训练中「路由均衡」不是自动发生的。
13.4 Mamba:线性时间的序列模型
Transformer 的注意力是平方复杂度,长序列(百万 token)成本爆炸。2023 年,Albert Gu 与 Tri Dao 提出 Mamba,基于状态空间模型(SSM),把序列建模变成线性递归:
hₜ = A·hₜ₋₁ + B·xₜ
yₜ = C·hₜMamba 的「选择性」体现在:矩阵 A、B、C 由当前输入动态决定,模型可以选择记住什么、遗忘什么——类似注意力,但保持线性复杂度。
examples/ch13_mamba.py 用一个 48 参数的极简 SSM 学习「累加」任务:
SSM 参数量:48
step 100:MSE = 0.88330
step 200:MSE = 0.11036
step 300:MSE = 0.06899
step 400:MSE = 0.05384
输入前 6 步: -0.0 0.5 -0.8 -0.7 -0.4 0.3
预测累加值: -0.01 0.58 -0.31 -0.96 -1.38 -1.15
真实累加值: -0.01 0.53 -0.29 -1.03 -1.42 -1.1548 个参数就学会了逐位累加(预测与真实几乎重合)。真实的 Mamba 用硬件友好的扫描算法并行化这个递归,在 2024 年成为 Transformer 的最强挑战者之一;Mamba-2、Jamba(混合 Mamba+注意力)进一步融合两条路线。
13.5 三条路线的启示
| 方法 | 解决什么 | 核心思想 | 代表 |
|---|---|---|---|
| LoRA | 微调成本 | 低秩增量 | QLoRA |
| MoE | 容量 vs 计算 | 稀疏专家路由 | Mixtral |
| Mamba | 长序列效率 | 选择性 SSM | Mamba-2 |
它们不是互斥的:今天的模型经常「LoRA 微调 + MoE 结构 + 长上下文技巧」一起用。
完整代码
本章用到的完整示例代码:
examples/ch13_lora.py
"""第 13 章示例(2023—2024):低秩适配(LoRA)微调。
冻结完整模型,只在权重旁加两个小矩阵做增量更新,大幅减少可训练参数。
任务:在 MNIST 上区分奇偶数。
运行方式:
uv run python examples/ch13_lora.py
"""
import torch
from torch import nn
from torch.utils.data import DataLoader, Subset
from torchvision import datasets, transforms
torch.manual_seed(0)
device = "cuda" if torch.cuda.is_available() else "cpu"
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,)),
])
train = datasets.MNIST("data", train=True, download=True, transform=transform)
train.targets = (train.targets % 2 == 0).long() # 奇偶二分类
loader = DataLoader(Subset(train, list(range(2000))), batch_size=128, shuffle=True)
test = datasets.MNIST("data", train=False, download=True, transform=transform)
test.targets = (test.targets % 2 == 0).long()
test_loader = DataLoader(test, batch_size=512)
class MLP(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(
nn.Flatten(),
nn.Linear(784, 64),
nn.ReLU(),
nn.Linear(64, 2),
)
def forward(self, x):
return self.net(x)
def evaluate(model):
model.eval()
correct = total = 0
with torch.no_grad():
for xb, yb in test_loader:
pred = model(xb.to(device)).argmax(1)
correct += (pred == yb.to(device)).sum().item()
total += len(yb)
return correct / total
def train(model, steps=200):
opt = torch.optim.Adam(model.parameters(), lr=1e-3)
for _ in range(steps):
xb, yb = next(iter(loader))
xb, yb = xb.to(device), yb.to(device)
opt.zero_grad()
loss = nn.functional.cross_entropy(model(xb), yb)
loss.backward()
opt.step()
# 方案 A:全量微调
model_full = MLP().to(device)
train(model_full)
print(f"全量微调:可训练参数 {sum(p.numel() for p in model_full.parameters()):,},"
f"测试准确率 {evaluate(model_full) * 100:.1f}%")
# 方案 B:LoRA 微调(冻结原权重,只训练低秩增量)
class LoRALayer(nn.Module):
def __init__(self, base, r=4, alpha=4):
super().__init__()
self.base = base
self.base.weight.requires_grad_(False)
self.base.bias.requires_grad_(False)
out_f, in_f = base.weight.shape # Linear 权重形状是 (out, in)
self.A = nn.Parameter(torch.randn(in_f, r) * 0.01)
self.B = nn.Parameter(torch.zeros(out_f, r))
self.scale = alpha / r
def forward(self, x):
return self.base(x) + (x @ self.A @ self.B.T) * self.scale
model_lora = MLP()
model_lora.net[1] = LoRALayer(model_lora.net[1])
model_lora = model_lora.to(device)
lora_params = sum(p.numel() for p in model_lora.parameters() if p.requires_grad)
train(model_lora)
print(f"LoRA 微调:可训练参数 {lora_params:,},测试准确率 {evaluate(model_lora) * 100:.1f}%")
lora_layer = model_lora.net[1]
delta_w = lora_layer.B @ lora_layer.A.T * lora_layer.scale
print(f"LoRA 增量矩阵 ||ΔW||/||W|| = {delta_w.norm().item() / lora_layer.base.weight.norm().item():.4f}")examples/ch13_moe.py
"""第 13 章示例(2023—2024):混合专家(MoE)。
多个"专家"子网络共享同一个路由器:每个输入只激活 top-2 个专家,
在几乎不增加计算量的前提下扩大模型容量。
运行方式:
uv run python examples/ch13_moe.py
"""
import torch
from torch import nn
torch.manual_seed(0)
device = "cuda" if torch.cuda.is_available() else "cpu"
# 合成二分类数据:两个交叠的高斯团
g = torch.Generator().manual_seed(0)
c0 = torch.randn(2000, 2, generator=g) * 0.8
c1 = torch.randn(2000, 2, generator=g) * 0.8 + torch.tensor([1.2, 0.4])
X = torch.cat([c0, c1])
y = torch.cat([torch.zeros(2000), torch.ones(2000)]).long()
class MoELayer(nn.Module):
def __init__(self, in_c, out_c, n_experts=4, top_k=2):
super().__init__()
self.experts = nn.ModuleList([
nn.Sequential(nn.Linear(in_c, 64), nn.ReLU(), nn.Linear(64, out_c))
for _ in range(n_experts)
])
self.router = nn.Linear(in_c, n_experts)
self.top_k, self.n_experts = top_k, n_experts
self.register_buffer("usage", torch.zeros(n_experts))
def forward(self, x):
logits = self.router(x)
topk = torch.topk(logits, self.top_k, dim=-1)
weights = topk.values.softmax(-1)
out = 0
for j, expert_idx in enumerate(topk.indices.T):
for i, e in enumerate(expert_idx):
self.usage[e.item()] += 1
out = out + self.experts[e](x) * weights[:, j:j + 1]
# 负载均衡辅助损失:让各专家使用频率接近
f = self.usage / max(self.usage.sum().item(), 1)
p = logits.softmax(-1).mean(0)
aux = (f * p).sum() * self.n_experts
return out, aux
class TinyMoE(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(nn.Linear(2, 32), nn.ReLU())
self.moe = MoELayer(32, 32)
self.head = nn.Linear(32, 2)
def forward(self, x):
h = self.net(x)
h, aux = self.moe(h)
return self.head(h), aux
model = TinyMoE().to(device)
opt = torch.optim.Adam(model.parameters(), lr=1e-3)
print(f"MoE 参数量:{sum(p.numel() for p in model.parameters()):,}")
for step in range(500):
idx = torch.randint(0, len(X), (128,), generator=g)
xb, yb = X[idx].to(device), y[idx].to(device)
opt.zero_grad()
logits, aux = model(xb)
loss = nn.functional.cross_entropy(logits, yb) + 0.1 * aux
loss.backward()
opt.step()
if (step + 1) % 100 == 0:
acc = (logits.argmax(1) == yb).float().mean().item()
print(f"step {step + 1}:loss = {loss.item():.4f},准确率 = {acc * 100:.1f}%")
usage = model.moe.usage.int().tolist()
print("各专家被激活次数:", usage, "→ 负载均衡损失让使用接近均匀")examples/ch13_mamba.py
"""第 13 章示例(2023—2024):简化版选择性状态空间模型(SSM)。
Mamba 的核心思想:状态转移参数由输入动态决定,
既能像 RNN 一样线性地处理长序列,又能像注意力一样"选择"要记住什么。
这里用一个极简 SSM 学习"累加"任务:预测输入序列的前缀和。
运行方式:
uv run python examples/ch13_mamba.py
"""
import torch
from torch import nn
torch.manual_seed(0)
device = "cuda" if torch.cuda.is_available() else "cpu"
T = 16
g = torch.Generator().manual_seed(0)
X = (torch.rand(2000, T, 1, generator=g) * 2 - 1)
y = X.cumsum(dim=1).clamp(-3, 3)
class SelectiveSSM(nn.Module):
"""h_t = a(x_t)·h_{t-1} + b(x_t)·x_t, y_t = c(x_t)·h_t。
参数 a/b/c 都由输入 x_t 决定——这就是"选择性"。"""
def __init__(self, d_state=8):
super().__init__()
self.a = nn.Linear(1, d_state)
self.b = nn.Linear(1, d_state)
self.c = nn.Linear(1, d_state)
def forward(self, x):
B, T, _ = x.shape
h = torch.zeros(B, 8, device=x.device)
outs = []
for t in range(T):
xt = x[:, t]
decay = torch.sigmoid(self.a(xt)) # 输入决定的"遗忘门"
h = decay * h + self.b(xt) * xt # 状态更新
outs.append((self.c(xt) * h).sum(-1, keepdim=True))
return torch.stack(outs, dim=1)
model = SelectiveSSM().to(device)
opt = torch.optim.Adam(model.parameters(), lr=1e-2)
print(f"SSM 参数量:{sum(p.numel() for p in model.parameters()):,}")
for step in range(400):
idx = torch.randint(0, len(X), (64,), generator=g)
xb, yb = X[idx].to(device), y[idx].to(device)
opt.zero_grad()
loss = nn.functional.mse_loss(model(xb), yb)
loss.backward()
opt.step()
if (step + 1) % 100 == 0:
print(f"step {step + 1}:MSE = {loss.item():.5f}")
model.eval()
with torch.no_grad():
pred = model(X[:1].to(device))[0, :6, 0]
print("输入前 6 步: ", " ".join(f"{v:.1f}" for v in X[0, :6, 0].tolist()))
print("预测累加值: ", " ".join(f"{v:.2f}" for v in pred.tolist()))
print("真实累加值: ", " ".join(f"{v:.2f}" for v in y[0, :6, 0].tolist()))动手实践
- 运行
examples/ch13_lora.py,把秩 r 从 4 改成 16,观察参数量、准确率与||ΔW||/||W||。 - 运行
examples/ch13_moe.py,把负载均衡损失权重从 0.1 改成 0(关闭),观察专家使用分布。 - 运行
examples/ch13_mamba.py,把状态维度从 8 改成 2,观察累加任务能否学会。 - 在 LoRA 实验里同时给两层都加 LoRA,对比单层 LoRA。
常见错误
| 错误写法 | 现象 | 原因 |
|---|---|---|
| LoRA 的 A/B 维度写反 | matmul 报错 | Linear 权重形状是 (out, in),A 是 (in, r),B 是 (out, r) |
LoRA 层替换后忘 .to(device) | device 不匹配 | 替换模块后要重新移动整个模型 |
| MoE 路由 softmax 维度错 | 权重不归一 | 在专家维度上 softmax |
| MoE 没有负载均衡损失 | 专家「塌缩」 | 路由器倾向少数专家,要加辅助损失 |
| Mamba 的递归用 Python 循环 | 长序列太慢 | 真实实现用并行扫描算法 |
| SSM 状态初始化不当 | 前期 loss 高 | h₀ 从 0 开始是标准做法 |
章末练习
基础
- LoRA 为什么把 ΔW 写成 B·Aᵀ?
- MoE 的 top-2 路由是什么意思?
- Mamba 的「选择性」体现在哪里?
提高
- 计算:784→64 的 Linear,全量微调 vs LoRA(r=4)的可训练参数各是多少。
- 在 MoE 示例中把 top_k 改成 1(硬路由),对比准确率与负载均衡。
- 在 Mamba 示例中把任务从「累加」改成「滑动平均」(yₜ = mean(xₜ₋₂..xₜ)),观察 SSM 能否学会。
挑战
- 实现 LoRA 的「合并」操作:训练后把 B·Aᵀ·scale 加回 W,验证合并后的模型与原模型输出一致。
- 实现简化 Mamba-2 风格的「块扫描」:把长度为 16 的序列分成 4 块并行递归,与逐步递归对比结果。
章末自测
- LoRA 中冻结的是?
- A. A、B 矩阵
- B. 原始权重 W
- C. 全部参数
- D. 偏置
- LoRA 的秩 r 通常取?
- A. 1000
- B. 4~64
- C. 与隐藏维度相同
- D. 1
- MoE 中决定 token 去哪个专家的是?
- A. 注意力
- B. 路由器
- C. 池化
- D. 卷积
- MoE 的负载均衡损失作用是?
- A. 提高精度
- B. 让专家使用均匀
- C. 加速训练
- D. 减少参数
- Mamba 的复杂度是?
- A. 平方
- B. 线性
- C. 指数
- D. 常数
- SSM 的递归公式是?
- A. hₜ = A·hₜ₋₁ + B·xₜ
- B. hₜ = xₜ
- C. hₜ = A·xₜ
- D. hₜ = hₜ₋₁
- Mixtral 是哪个机构的 MoE 模型?
- A. OpenAI
- B. Mistral AI
- C. Google
- D. Meta
- LoRA 微调 70B 模型的关键优势是?
- A. 精度更高
- B. 可训练参数极少
- C. 推理更快
- D. 数据更少
- QLoRA 在 LoRA 基础上加了?
- A. 量化
- B. 蒸馏
- C. 剪枝
- D. 集成
- 2024 年主流开源 MoE 模型的专家激活方式是?
- A. 全部激活
- B. top-2 稀疏激活
- C. 随机激活
- D. 不激活
