Skip to content

第 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ᵀ)·scale

A 是 (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.4058

LoRA 只训练 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.15

48 个参数就学会了逐位累加(预测与真实几乎重合)。真实的 Mamba 用硬件友好的扫描算法并行化这个递归,在 2024 年成为 Transformer 的最强挑战者之一;Mamba-2、Jamba(混合 Mamba+注意力)进一步融合两条路线。

13.5 三条路线的启示 ​

方法解决什么核心思想代表
LoRA微调成本低秩增量QLoRA
MoE容量 vs 计算稀疏专家路由Mixtral
Mamba长序列效率选择性 SSMMamba-2

它们不是互斥的:今天的模型经常「LoRA 微调 + MoE 结构 + 长上下文技巧」一起用。

完整代码 ​

本章用到的完整示例代码:

examples/ch13_lora.py ​

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 ​

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 ​

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()))

动手实践 ​

  1. 运行 examples/ch13_lora.py,把秩 r 从 4 改成 16,观察参数量、准确率与 ||ΔW||/||W||。
  2. 运行 examples/ch13_moe.py,把负载均衡损失权重从 0.1 改成 0(关闭),观察专家使用分布。
  3. 运行 examples/ch13_mamba.py,把状态维度从 8 改成 2,观察累加任务能否学会。
  4. 在 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 开始是标准做法

章末练习 ​

基础

  1. LoRA 为什么把 ΔW 写成 B·Aᵀ?
  2. MoE 的 top-2 路由是什么意思?
  3. Mamba 的「选择性」体现在哪里?

提高

  1. 计算:784→64 的 Linear,全量微调 vs LoRA(r=4)的可训练参数各是多少。
  2. 在 MoE 示例中把 top_k 改成 1(硬路由),对比准确率与负载均衡。
  3. 在 Mamba 示例中把任务从「累加」改成「滑动平均」(yₜ = mean(xₜ₋₂..xₜ)),观察 SSM 能否学会。

挑战

  1. 实现 LoRA 的「合并」操作:训练后把 B·Aᵀ·scale 加回 W,验证合并后的模型与原模型输出一致。
  2. 实现简化 Mamba-2 风格的「块扫描」:把长度为 16 的序列分成 4 块并行递归,与逐步递归对比结果。

章末自测 ​

  1. LoRA 中冻结的是?
    • A. A、B 矩阵
    • B. 原始权重 W
    • C. 全部参数
    • D. 偏置
  2. LoRA 的秩 r 通常取?
    • A. 1000
    • B. 4~64
    • C. 与隐藏维度相同
    • D. 1
  3. MoE 中决定 token 去哪个专家的是?
    • A. 注意力
    • B. 路由器
    • C. 池化
    • D. 卷积
  4. MoE 的负载均衡损失作用是?
    • A. 提高精度
    • B. 让专家使用均匀
    • C. 加速训练
    • D. 减少参数
  5. Mamba 的复杂度是?
    • A. 平方
    • B. 线性
    • C. 指数
    • D. 常数
  6. SSM 的递归公式是?
    • A. hₜ = A·hₜ₋₁ + B·xₜ
    • B. hₜ = xₜ
    • C. hₜ = A·xₜ
    • D. hₜ = hₜ₋₁
  7. Mixtral 是哪个机构的 MoE 模型?
    • A. OpenAI
    • B. Mistral AI
    • C. Google
    • D. Meta
  8. LoRA 微调 70B 模型的关键优势是?
    • A. 精度更高
    • B. 可训练参数极少
    • C. 推理更快
    • D. 数据更少
  9. QLoRA 在 LoRA 基础上加了?
    • A. 量化
    • B. 蒸馏
    • C. 剪枝
    • D. 集成
  10. 2024 年主流开源 MoE 模型的专家激活方式是?
    • A. 全部激活
    • B. top-2 稀疏激活
    • C. 随机激活
    • D. 不激活