Skip to content

WARNING

🧪 Beta公测版本提示:教程主体已完成,正在优化细节,欢迎大家提Issue反馈问题或建议。

混合专家 MoE — demo.py 代码详解

Download demo.py

运行方式

bash
cd docs/nn-decision/dl/moe/code
python demo.py

CPU + NumPy 即可,几十秒。四个高斯簇上的二分类:线性路由器 + 4 个线性专家,每个样本只唤醒 Top-2。同一数据训两遍(有/无负载均衡),画出专家使用率、分类损失、决策边界。没有 Transformer、没有专家并行,只把门控、稀疏加权和 Laux 跑通。

代码逐段详解

第1步:导入与超参 — 每个名字后面干什么

python
N_EXPERTS = 4
TOP_K = 2
LR = 0.08
STEPS = 400
AUX_COEF = 0.05
  • os / _IMAGES_DIRos.path.dirname(os.path.abspath(__file__)) 是当前 .py 所在目录,再 join(..., '..', 'images'),无论从哪启动路径都对。
  • numpy:造数据、矩阵乘、手写梯度。本章不用 PyTorch,为的是把 Softmax 门控和辅助损失摊成数组。
  • font.sans-serif:中文字体列表,缺第一个就试下一个。axes.unicode_minus = False:否则负号画成方块。
  • np.random.seed(42):簇采样和 W_r/W_e 初始化共用这一份随机源。
  • N_EXPERTS=4:和四个簇对齐,鼓励「一块区域一个专家」。标签却只有 0/1 交替,所以专家学的是区域,不是四分类。
  • TOP_K=2:每个样本激活 2/4 专家,对应 Mixtral 一类的稀疏门控。
  • HIDDEN=8:声明了,但线性专家没用到——不要按名字脑补一层 MLP。
  • AUX_COEF:正文里的 α。太大则路由被「均匀」绑架;太小则塌缩。

第2步:make_data — 为什么用四簇而不是一团云

路由器要学「不同区域走不同专家」。四个中心分居象限,簇内噪声 0.35,彼此不太糊在一起。

python
centers = np.array([[-1.5, -1.2], [1.6, -1.0], [-1.4, 1.5], [1.5, 1.4]])
for i, c in enumerate(centers):
    xs.append(c + 0.35 * np.random.randn(n_per, 2))
    ys.append(np.full(n_per, i % 2))
  • i % 2:簇 0、2 → 类 0,簇 1、3 → 类 1。对角同类,单靠一个线性分类器吃力;MoE 可以让两个专家各管一块。
  • np.full(n_per, i % 2):长度 n_per、值全相同的标签数组。
  • np.random.permutation(len(X)):打乱下标再切片。astype(float) 让标签能进 BCE 的 (1-y)*log(1-p)

第3步:softmax / sigmoid — 门控概率和分类概率

路由器输出未归一化 logits,门控是

gi(x)=ezijezj,z=xWr+br
python
z = logits - logits.max(axis=axis, keepdims=True)
e = np.exp(z)
return e / e.sum(axis=axis, keepdims=True)
  • 先减 maxsoftmax(z)=softmax(zc)。某个 z 很大时不减,exp 会溢出成 inf,整行变 NaN。
  • keepdims=Truemax 后仍留着被缩掉的那一维,才能和 logits 广播相减。axis=1(B,4)(B,1),不是 (B,)
  • axis=-1:默认沿最后一维;对 (B, n_experts) 就是对专家维做。
python
return 1.0 / (1.0 + np.exp(-np.clip(x, -40, 40)))

专家加权和是标量 logit,分类概率 p=σ(y^)clip±40:挡住 exp 溢出,和 Softmax 减 max 是同一类数值卫生。


第4步:TinyMoE 两套权重

python
self.W_r = np.random.randn(in_dim, n_experts) * scale   # (2, 4)
self.W_e = np.random.randn(n_experts, in_dim) * scale   # (4, 2)
符号代码形状角色
Wr,brW_r,b_r(2,4) / (4,)路由器:每样本一个 4 维 logit
We(i)W_e[i]每行 (2,)i 个专家:Ei(x)=wix+bi

专家也做成线性,是因为本章要看的是路由,不是专家容量。四个超平面靠门控拼出一块块决策。这不是 nn.Module,没有 super().__init__(),更新靠手写 SGD。


第5步:route — Softmax → Top-2 → 再归一化

g~=renorm(Top-k(g))
python
logits = X @ self.W_r + self.b_r
probs = softmax(logits, axis=1)
top_idx = np.argsort(probs, axis=1)[:, -TOP_K:]
rows = np.arange(len(X))[:, None]
top_p = probs[rows, top_idx]
top_p = top_p / (top_p.sum(axis=1, keepdims=True) + 1e-9)
  • X @ W_r(B,2)@(2,4)→(B,4)@ 是矩阵乘;* 才是逐元素,这里不能混。
  • np.argsort(..., axis=1):每行从小到大的下标[:, -2:] 取最后两列 = 概率最大的两个专家。: 是「这一维全要」,-2: 是「从倒数第二个到末尾」。
  • np.arange(B)[:, None](B,) 加成 (B,1),才能和 top_idx(B,2) 一起做高级索引。
  • probs[rows, top_idx]:第 b 行取出那两个专家的 g,得到 (B,k)
  • 再除以和:丢掉的专家不参与加权。+1e-9 防止除零。

未选中的专家本步不进 y^。实现上 expert_out 仍一次算出全部 4 列再切片——batch 很小,不是生产级稀疏内核。


第6步:forward — 只把被选专家加权求和

python
eo = self.expert_out(X)       # (B, 4)
chosen = eo[rows, top_idx]    # (B, k)
y_hat = (chosen * top_p).sum(axis=1)
y^=j=1kg~ijEij(x)
  • expert_outX @ W_e.T + b_eW_e(4,2).T 转成 (2,4),一次得到每个专家的标量。
  • chosen * top_p:逐元素。第 j 个被选输出乘它的门控,再沿 axis=1 加总成一个 logit。
  • 没进 Top-2 的专家对 y^ 没有前向贡献;train_step 也只给选中的 e 累加梯度。

第7步:load_balance_loss — 为什么要 NfiPi

放任不管,路由器会把票永远投给最先碰巧有用的一两个专家,其余 W_e 收不到梯度(路由崩溃)。Switch 一类写法:

Laux=Ni=1NfiPi
python
for k in range(TOP_K):
    for i in top_idx[:, k]:
        f[i] += 1.0
f = f / (B * TOP_K)
P = probs.mean(axis=0)
return float(self.n * np.sum(f * P))
  • fi:专家 iB×k 次选择里被点到的频率。每个样本贡献 k 次,所以除以 B * TOP_Ktop_idx[:, k] 是「所有样本的第 k 个被选下标」。
  • Pi:Top-k 之前的 Softmax 在 batch 上平均。用完整 g 而不是 g~,鼓励选之前就把质量摊开。
  • 都均匀fi=Pi=1/N,乘 N1;塌缩到一个专家则 fP 都尖,乘积变大。
  • float(...):后面要和 Python 标量 CE 相加、还要 print

第8步:train_step — BCE、专家梯度、路由器直通

python
p = sigmoid(y_hat)
loss = float(-np.mean(y * np.log(p + eps) + (1 - y) * np.log(1 - p + eps)))
aux = self.load_balance_loss(probs, top_idx) if use_aux else 0.0
total = loss + (AUX_COEF * aux if use_aux else 0.0)
dlogit = (p - y) / len(y)

二元交叉熵。p=σ(y^)/y^=py(再除 B 做平均)。eps=1e-9:挡住 log(0)use_aux=Falseaux 记 0,分类梯度照常,只是不加压均衡。

专家:只更新被选中的行。

python
e = top_idx[b, j]
w = top_p[b, j]
dW_e[e] += dlogit[b] * w * X[b]
db_e[e] += dlogit[b] * w

y^Ee 乘了门控 w,链式法则把 dy^w 再乘 x。没被选的 edW_e[e] 保持 0。

路由器:Top-k 离散、不可导。 代码用直通近似——把「这个专家的输出 × 误差」当作对 ge 的信号:

python
d_probs[b, e] += dlogit[b] * eo[b, e]
d_logits = d_probs - (d_probs * probs).sum(axis=1, keepdims=True) * probs

Ee 和误差同号,就希望提高 ge。第二行是 Softmax 雅可比的粗糙近似(源码注释写明了):精确式是 g(v1(gv)),这里略去最外层乘 g,够推动路由,不当精确反传教材。

有辅助损失时再推 P

python
dP = AUX_COEF * self.n * f / len(X)
d_logits += dP

Laux/PiNfi,广播到每个样本的 logits。这是轻推均匀,不是把分类梯度关掉。

python
self.W_e -= LR * dW_e
self.W_r -= LR * (X.T @ d_logits)

手写 SGD。X.T @ d_logits(2,B)@(B,4)→(2,4)b_rd_logits.mean(axis=0):偏置对 batch 平均。


第9步:train / main — 对照实验在验证什么

train(use_aux, X, y) 每步吃全量 X(没有再切 mini-batch)。expert_usagetop_idx 里出现的次数除以总和,和 f 同一件事。每 50 步打印 usage

maintrain(True)train(False)。两次各自 TinyMoE(),不要共用权重。

  1. 专家负载柱:无均衡往往一根特别高;有均衡四根更近。这是 Laux 存在的理由。
  2. 分类 CE:均衡不该把分类训废。两条都应下降;有时有均衡略高,那是 α 的税。
  3. 决策边界meshgrid + np.c_[xx.ravel(), yy.ravel()] 铺成 (40000,2),用有均衡模型的 sigmoid(forward) 填色。np.c_ 按列拼接两个展平坐标。

关键概念速查表

概念数学 / 直觉代码
门控g=softmax(xWr+br)routesoftmax(logits)
Top-2只点亮最大 k 个,再归一化argsort + [:, -TOP_K:]
稀疏输出y^=jTkg~jEj(chosen * top_p).sum(1)
辅助损失NfiPi,防塌缩load_balance_loss
BCE[ylogp+(1y)log(1p)]train_steploss
专家梯度只回传到被选 edW_e[e] += dlogit * w * x
Softmax 反传粗糙:dzdg()gd_logits = d_probs - ...
keepdims / [:, None]广播对齐softmax、高级索引

源码位置

clone 后打开(相对仓库根目录):

docs/nn-decision/dl/moe/code/demo.py