第25章 双头模型:共享底座 + 两个分支
小白一句话
两个目标(回访、持续)要训练两个东西。最省事的做法不是各训各的,而是让它们共用一层底座------底座学"这个人什么样",上面再长两个分支,一个答回访、一个答持续。这就是双头模型。
这一章在干嘛
第24章造好了双标签(returned、sustained_active_7d)。这一章回答"怎么训练":两个目标需要两个模型吗?还是可以共用一个?答案是后者,但"共用"在树模型世界里长什么样,得说清楚。
顺带把第24章埋的那个关键概念落地:条件头只在回访过的样本上训练 ------损失公式里的 return 掩码。
两个方案,选哪个
方案 A:训两个完全独立的模型。 一个 GBDT 学 returned(全量 105 行),另一个 GBDT 学 sustained(全量 105 行)。各训各的,互不干扰。
方案 B:一个双头模型。 一个特征底座(叫 trunk),上面挂两个头------回访头、条件持续头。头之间共享底座的表示。
方案 B 为什么好,两个实在的理由:
- 共享底座 = 表示复用。两个目标吃的是同一份特征、同一个"这个人什么样"的中间表示。回访头要的"新近度、趋势"信号,条件头要的"活跃率、持续性"信号,都在同一份特征里,底座一次学完,两个头各取所需,不用各学一遍。
- 省一次链路。特征工程、切分、标准化、权重体系,都只做一份。两个头共用。
代价也直说:两个目标共享底座后,训练时彼此的梯度会互相牵扯(一个头想往这边调,另一个想往那边),需要调 λ 平衡------后面损失公式会讲。
树模型世界里的"双头"
字面意义的 trunk + heads,在神经网络里是自然的长法(第15章那个 MLP,最后一层拆成两个输出就成)。但本项目的主模型是 GBDT------树模型没有现成的双头,加法模型一次只能输出一个标量。
所以在树模型世界里,"双头"的实际落地是:
css
特征底座(10 维特征 + 一套切分 + 一套权重)
┌──────────────────────────────┐
│ │
回访头(GBDT,全量样本训练) 条件头(GBDT,只吃 returned=1 的样本)
学 P(回访) 学 P(持续 | 回访)
两个 GBDT,吃同一份特征和切分,各答各的------"共享"体现在特征、切分、权重体系、超参都只做一份,而不是共享模型参数本身。这在工程上的收益和神经网络的共享底座是一样的:链路只维护一条。
落到代码就是两行 fit,掩码体现在条件头的训练子集下标(完整脚本 附件/07_双目标模型升级篇/activity_dual_head.py):
python
# 双标签与切分(口径同第 24 章):X_tr 105 行,ret_tr/sus_tr 两个标签
mask_tr = ret_tr == 1 # 条件头的训练子集:只留回访过的 92 行
# 回访头:全量训练,学 P(回访)
gb_ret = GradientBoostingClassifier(random_state=20260827, n_estimators=100,
max_depth=3, learning_rate=0.1)
gb_ret.fit(X_tr, ret_tr)
# 条件头:只在 returned=1 的子集上训练,学 P(持续 | 回访)------掩码就是这行
gb_sus = GradientBoostingClassifier(random_state=20260827, n_estimators=100,
max_depth=3, learning_rate=0.1)
gb_sus.fit(X_tr[mask_tr], sus_tr[mask_tr])
mask_tr = ret_tr == 1 这一行,就是第 24 章那个"条件概率只在回访前提下有定义"的代码形态。两个模型喂的是同一份 X_tr(特征、切分、超参一份),只有条件头多了一道子集下标------这就是树模型世界里的双头。
神经网络版的双头:trunk + heads 的字面长法
树模型的"共享"是数据链路层面的。神经网络里双头是字面意义的:一个隐层当底座,上面长两个输出头,底座参数真真切切只有一份。第 12 章讲过神经网络在做什么(加权求和 + 压扁),这里把它从"一个输出头"改成"两个输出头",就是双头:
css
输入 10 维 ──► 共享隐层 32(ReLU)──► 回访头(sigmoid)→ P(回访)
└────► 条件头(sigmoid)→ P(持续 | 回访)
第 15 章那个 MLP 最后一层拆成两个输出,就是这个结构。和树模型版一一对应的地方:
- 共享底座 :
W1/b1只有一份,两个头共用。反向传播时,回访头的梯度和条件头的梯度在共享层相加,一起更新这一份底座------"表示复用"是物理发生的,不是概念上的。 - 掩码 :条件头的 BCE 损失只在
returned=1的训练子集(92 行)上算,没回访的 13 行不参与条件头损失(但照常参与回访头损失)。 - λ 平衡 :总损失 =
L(回访头, 全量) + λ × L(条件头, 子集),λ 默认 1。
核心代码(numpy 手写,不引 torch;完整版 附件/07_双目标模型升级篇/activity_dual_nn.py):
python
def train_dual(X, ret, sus, mask, H=32, lr=0.02, epochs=3000, l2=1e-4):
W1, b1, Wr, br, Ws, bs = init(H) # 小随机初始化
retc = ret.reshape(-1, 1) # 标签统一列向量
Xm = X[mask] # 条件头只用回访子集
susc = sus[mask].reshape(-1, 1)
for ep in range(epochs):
# 前向:共享底座算一次,两个头各算各的
h = relu(X @ W1 + b1) # 共享隐层,一份 W1/b1
pr = sigmoid(h @ Wr + br) # 回访头,全量
ps = sigmoid(relu(Xm @ W1 + b1) @ Ws + bs) # 条件头,只在子集上算
# 反向:条件头的梯度只放进子集行,其余为 0
dzr = (pr - retc) / n
dzs = zeros((n, 1)); dzs[mask] = (ps - susc) / n_m
# 共享底座梯度 = 两个头的梯度相加,一起反传
dH = (dzr @ Wr.T + dzs @ Ws.T) * (h > 0)
W1 -= lr * (X.T @ dH + l2 * W1) # 底座被两个头共同更新
...
逐段看这几行在干什么:
- 前向一次 :
h = relu(X @ W1 + b1)只算一遍底座,两个头从同一个h各取所需------回访头全量算,条件头用X[mask]的子集算。掩码在神经网络里不是"跳过损失",而是条件头的前向输入就是子集。 - 梯度分叉 :
dzr是全量行的梯度;dzs是个零向量,只在mask为真的行上被填上条件头的梯度------没回访的 13 行对条件头没有梯度,自然不更新条件头。 - 共享层合流 :
dH = dzr @ Wr.T + dzs @ Ws.T------回访头和条件头的梯度在底座这里相加 ,一份W1同时被两个目标拉着走,这就是"梯度互相牵扯、要调 λ"的物理位置。
在这份数据上,神经网络版和树模型版各跑一遍(同切分、同评估口径,脚本都在本章配套里):
| 实现 | 回访头(全量 57 人) | 条件头(回访 42 人) |
|---|---|---|
| 树模型(两个 GBDT) | PR-AUC 0.827 / AUC 0.689 | PR-AUC 0.864 / AUC 0.814 |
| 神经网络(共享底座双头) | PR-AUC 0.899 / AUC 0.748 | PR-AUC 0.912 / AUC 0.884 |
| 神经网络换 3 个种子 | 0.898 ~ 0.902 | 0.903 ~ 0.919 |
两个诚实的地方要摆出来:
- NN 分数略高,但那是噪声级。回访头 0.899 vs 0.827、条件头 0.912 vs 0.864,差 0.03~0.07------57/42 人的测试集上,这个差距和第 21 章 bootstrap 说的"小样本分不出"是同一件事。别拿这张表去给"NN 更好"当证据。
- NN 这次换种子很稳(0.898~0.902),因为它用了全量梯度下降 + 小学习率 + L2 正则,和第 15 章那个抖的 sklearn MLP 不是一种训法。但代价是这些超参(学习率、轮数、L2、初始化尺度)是调出来的------树模型开箱即用,NN 每换一份数据都得重新调。
那选哪个?工程上还是树模型版:不调参、能出特征重要性(第 22 章)、和全书其它章节同一套链路口径。神经网络版的价值在概念------它是"共享底座"这三个字最不含糊的实物,看懂了它,第 26 章"两个头乘起来"、第 27 章"每个头单独评"就都有了具象。
条件头的掩码:只在回访子集上训练
第24章说过,sustained 只在"回访过"的人身上有定义------没回访的人,谈"持不持续"没有意义。所以条件头的训练要带一个 return 掩码:损失只在这部分样本上算。
scss
总损失 = L(回访头, 全部样本) + λ × L(条件头, 只在 returned=1 的样本上算)
λ 是平衡两个头权重的旋钮(默认 1,谁更重要可以调)。树模型落地时,掩码就体现在"条件头只用训练集里 returned=1 的那 92 行来训"------比全量 105 行少了 13 个"没回访"的样本。
那 13 个样本为什么不进条件头?做一次对照就明白。同样在测试集 42 个回访用户上评估:
| 条件头怎么训的 | 学的是什么 | PR-AUC | ROC-AUC |
|---|---|---|---|
| B 只在回访子集训(掩码) | P(持续 | 回访),口径一致 | 0.864 | 0.814 |
| C 全量训(不掩码) | P(持续),把没回访的当负例 | 0.884 | 0.840 |
单看数字,C 还略高一点。别急着下"不掩码更好"的结论------42 人的测试子集,0.02 的差在噪声里,分不出高下。真正的区别在目标函数:
- B 学的是条件概率 P(持续 | 回访),评估也在回访子集上做------口径一致;
- C 学的是全量边缘概率 P(持续),训练时把"没回访"的人也当成了"不持续"的负例,但评估却只看回访子集------口径对不上。它偷了 13 个训练样本,换来的却是"猜错对象"。
真实项目里测试集会有成千上万的回访用户,混入子集外样本造成的偏差会被放大。所以掩码不是可有可无的优化,是条件概率的定义要求------没回访的人,条件概率根本没定义。
三个头的完整数字
| 头 | 训练口径 | 评估子集 | PR-AUC | ROC-AUC |
|---|---|---|---|---|
| A 回访头 | 全量 105 行,学 returned | 全部 57 人 | 0.827 | 0.689 |
| B 条件头 | 回访 92 行(掩码),学 sustained | 回访 42 人 | 0.864 | 0.814 |
回访头 0.827 和第15章那个"全特征 PR-AUC 0.827"是同一个数------回访头就是单目标模型原封不动;条件头 0.864 是新增的能力。两个数各管一摊:回访头决定"要不要触达",条件头决定"值不值得深耕"。
先说清口径,免得后面章节对不上号:这张表的两个数都是未校准的 GBDT 原始输出------本章只关心"掩码训练 vs 不掩码"的目标函数差异,还没做校准。第27章会在校准后重评(条件头变成 0.896),数字不同是校准造成的,不是谁算错了。
(神经网络版的对应数字在上一节那张对比表里:回访 0.899、条件 0.912,分数略高但噪声级,选型不靠它。)
它和前后章节的关系
- 第24章造了双标签,本章把它们接上模型;
- 第26章把两个头的输出拼成对外概率:
P(持续) = P(回访) × P(持续 | 回访)------两个头各出各的,乘起来用; - 第27章评估时,回访头在全量上评、条件头只在回访子集上评(本章已经按这个口径在数了);
- 双头共享底座的思路,和第15章那个"换个模型试试"里的 MLP 是同一个思路------那里是换模型,这里是加目标;神经网络版的双头就是那个 MLP 拆成两个输出头,本章给了它的 numpy 实现。
动手
- 跑
附件/07_双目标模型升级篇/activity_dual_head.py,对照三个头的数字:回访头 0.827/0.689、条件头 0.864/0.814。 - 把脚本里条件头的训练从"子集"改成"全量"(去掉
[mask_tr]那两处下标),重跑看 B 和 C 的数字------观察 42 人子集上差多少,看看"口径不一致但数字分不出"是什么感觉。 - 把
λ换成 0.5 或 2 的想法过一遍:条件头权重调低,模型会更偏向把回访答准;调高则更偏向持续。真实项目里 λ 怎么定?先在脚本里把条件头换成sample_weight乘 2 跑一版,看条件头 PR-AUC 动没动。 - 跑
附件/07_双目标模型升级篇/activity_dual_nn.py,对照神经网络版的两个数字(回访 0.899/0.748、条件 0.912/0.884)。把H从 32 改成 64 重跑,看两个头变多少------隐层越大参数越多,105 行样本喂得饱吗? - 在
activity_dual_nn.py里把dzs那行的掩码去掉(改成dzs = (ps_full - sus_full) / n,条件头用全量算),重跑对比------不掩码的神经网络版条件头,和树模型版 C 是同一个错误。
本章配套脚本
附件/07_双目标模型升级篇/activity_dual_head.py(回访头全量训练 + 条件头掩码子集训练 + 不掩码对照,同切分同超参)与附件/07_双目标模型升级篇/activity_dual_nn.py(numpy 手写共享底座双头 MLP:前向一次、条件头子集损失、共享层梯度相加,换种子稳定性对比,不引 torch)。用 scikit-learn 与 numpy,已在第4章附件/00_公共/requirements.txt列出。双标签口径与第24章一致,训练表由第12章生成。