Wukong: Towards a Scaling Law for Large-Scale Recommendation
-
- 研究动机与背景
- 一、整体架构概览
- [二、Wukong Layer ------ 核心公式推导](#二、Wukong Layer —— 核心公式推导)
-
- [2.1 Factorization Machine Block (FMB)](#2.1 Factorization Machine Block (FMB))
- [2.2 Linear Compression Block (LCB)](#2.2 Linear Compression Block (LCB))
- [2.3 合并 + 残差 + LayerNorm](#2.3 合并 + 残差 + LayerNorm)
- [2.4 交互阶数为何指数增长?](#2.4 交互阶数为何指数增长?)
- [三、Scaling Strategy(Dense Scaling)](#三、Scaling Strategy(Dense Scaling))
- [四、PyTorch 简化实现](#四、PyTorch 简化实现)
- 五、关键点回顾
《Wukong: Towards a Scaling Law for Large-Scale Recommendation》(Meta, ICML 2024) 的核心内容是用堆叠因子分解机(Stacked FM)构建可"稠密扩展(Dense Scaling)"的推荐模型,并首次在推荐域验证跨两个数量级的 Scaling Law 。
论文原文:https://arxiv.org/abs/2403.02545
研究动机与背景
近年来,Scaling Law(缩放定律)已成为大语言模型(LLM)持续提升性能的重要基础------模型质量随计算量、参数量和训练数据量的增加而呈现出可预测的幂律提升。然而,深度学习推荐系统(DLRS)长期以来未能建立起类似的 Scaling Law,其根本原因在于传统推荐模型的扩展机制存在结构性缺陷。
传统推荐模型普遍依赖"Sparse Scaling(稀疏扩展)",即通过扩大 Embedding Table 的行数(特征 ID 数)和维度来增加参数量。 这种扩展方式存在两大核心问题:
- 无法增强高阶特征交互能力:单纯扩大稀疏嵌入参数并不能提升模型捕捉特征间复杂高阶交互的能力,随着嵌入表增大,收益迅速饱和,甚至出现边际效用递减。
- 硬件利用率低下:Embedding Lookup 属于访存密集型(Memory-bound)操作,受限于内存带宽而非计算能力,无法利用现代 AI 加速器(如 GPU/TPU)日益增长的浮点算力。这导致单纯扩大嵌入表带来极高的基础设施成本,却难以转化为实际训练吞吐量和效果的提升。
与此同时,工业级推荐场景的数据规模呈指数级增长(生产数据集可达数百亿至千亿级样本),对模型容量和表达能力的诉求日益迫切。现有主流推荐架构(如 DLRM、DCNv2、xDeepFM 等)或因仅支持有限阶特征交互,或因缺乏系统化的稠密扩展方案,在参数量和计算量增大时往往快速趋于饱和甚至性能退化,无法满足随数据/算力持续增长而稳定提升模型质量的诉求。
为此,Meta 研究团队提出 Wukong ,旨在探索推荐领域的 Scaling Law------其核心动机是摒弃单一依赖 Sparse Scaling 的旧范式,转而通过Dense Scaling(稠密扩展),即系统化扩大特征交互组件(Interaction Component)的计算量(如堆叠更深/更宽的交互层),使推荐模型能够像 LLM 一样在跨数量级的模型复杂度(>100 GFLOP/example,等效 GPT-3 级训练算力)上持续、可预测地降低损失,从而验证推荐域 Scaling Law 的存在性。
一、整体架构概览
Sparse Features ──► Embedding Lookup ──┐
├──► X₀ ∈ ℝ^{n×d}
Dense Features ───► MLP Projection ─────┘
│
▼
┌─── Interaction Stack ───┐
│ l × Wukong Layer │
│ (FMB + LCB + Res + LN) │
└─────────────────────────┘
│
▼
Final MLP ──► Sigmoid ──► CTR Prediction
- 输入:类别型特征经 Embedding Lookup,连续型特征经 MLP 投影,统一维度 d,组成矩阵 X₀ ∈ ℝ^{n×d}(n = 特征嵌入总数)。
- 核心:Interaction Stack 堆叠 l 个相同的 Wukong Layer,逐层以"类二进制幂"方式使可捕获的特征交互阶数翻倍(第 i 层理论最高捕获 2ⁱ 阶交互)。
- 输出:Interaction Stack 结果展平后送入 Final MLP,做 CTR 预测。
二、Wukong Layer ------ 核心公式推导
第 i 层输入记为 Xᵢ ∈ ℝ^{nᵢ×d} ,输出 Xᵢ₊₁ ∈ ℝ^{nᵢ₊₁×d}。每层含两个并行分支:
2.1 Factorization Machine Block (FMB)
FMB 负责显式二阶交叉,公式如下:
FM ( X i ) = X i X i ⊤ ∈ R n i × n i \text{FM}(X_i) = X_i X_i^\top \in \mathbb{R}^{n_i \times n_i} FM(Xi)=XiXi⊤∈Rni×ni
- 朴素 FM 内积矩阵 XᵢXᵢᵀ 复杂度为 O(nᵢ²d),存储 O(nᵢ²)。
- 低秩优化 :当 d ≪ n 时,XᵢXᵢᵀ 是低秩矩阵。引入可学习投影矩阵 Y ∈ ℝ^{nᵢ×k} (k ≪ nᵢ),用结合律先算 XᵢᵀY ∈ ℝ^{d×k},再算 Xᵢ(XᵢᵀY) ∈ ℝ^{nᵢ×k},复杂度降为 O(nᵢdk):
FM eff ( X i ) = X i ( X i ⊤ Y ) ∈ R n i × k \text{FM}_{\text{eff}}(X_i) = X_i (X_i^\top Y) \in \mathbb{R}^{n_i \times k} FMeff(Xi)=Xi(Xi⊤Y)∈Rni×k
- 随后 Flatten → LayerNorm → MLP → Reshape 回嵌入形式:
FMB ( X i ) = reshape ( MLP ( LN ( flatten ( FM eff ( X i ) ) ) ) ) ∈ R n F × d \text{FMB}(X_i) = \text{reshape}\Big(\text{MLP}\big(\text{LN}(\text{flatten}(\text{FM}_{\text{eff}}(X_i)))\big)\Big) \in \mathbb{R}^{n_F \times d} FMB(Xi)=reshape(MLP(LN(flatten(FMeff(Xi)))))∈RnF×d
其中 n_F 是 FMB 输出的嵌入个数(超参数)。
2.2 Linear Compression Block (LCB)
LCB 只做线性重组,不增加交互阶数,作用是保留原始一阶信息并控制宽度:
LCB ( X i ) = W L X i ∈ R n L × d , W L ∈ R n L × n i \text{LCB}(X_i) = W_L X_i \in \mathbb{R}^{n_L \times d}, \quad W_L \in \mathbb{R}^{n_L \times n_i} LCB(Xi)=WLXi∈RnL×d,WL∈RnL×ni
n_L 是 LCB 输出嵌入个数。
2.3 合并 + 残差 + LayerNorm
FMB 与 LCB 输出在特征维拼接,再与输入 Xᵢ 做残差连接,最后 LayerNorm:
X i + 1 = LN ( concat ( FMB ( X i ) , LCB ( X i ) ) + X i ) ∈ R ( n F + n L ) × d X_{i+1} = \text{LN}\Big(\text{concat}\big(\text{FMB}(X_i),\ \text{LCB}(X_i)\big) + X_i\Big) \in \mathbb{R}^{(n_F+n_L) \times d} Xi+1=LN(concat(FMB(Xi), LCB(Xi))+Xi)∈R(nF+nL)×d
注:论文原公式写作
LN(concat(FMB, LCB) + Xᵢ),实际实现中残差连接作用于拼接结果与输入之间,LayerNorm 稳定训练。
2.4 交互阶数为何指数增长?
- 归纳假设:第 i 层输入含 1 ~ 2ⁱ⁻¹ 阶交互。
- FMB 做两两二阶交叉:阶数 o₁ + o₂,最大 2ⁱ⁻¹ + 2ⁱ⁻¹ = 2ⁱ。
- LCB 保持 1 ~ 2ⁱ⁻¹ 阶不变。
- 拼接后第 i+1 层输入含 1 ~ 2ⁱ 阶交互 → 堆叠 l 层最高捕获 2ˡ 阶特征交叉。
三、Scaling Strategy(Dense Scaling)
Wukong 采用协同扩展(Synergistic Upscaling),优先放大交互相关参数:
| 超参数 | 含义 | 扩展优先级 |
|---|---|---|
| l | Interaction Stack 层数 | ⭐⭐⭐ 最高(先加层) |
| n_F | FMB 输出嵌入数 | ⭐⭐ |
| n_L | LCB 输出嵌入数 | ⭐ |
| k | FM 低秩投影维度 | ⭐⭐ |
| MLP 深度/宽度 | FMB 内 MLP 容量 | ⭐⭐ |
论文在 Meta 内部 146B 样本数据集上将 Dense 参数量从 0.74B 扩到 17B,GFLOP/example 从 1 扩到 >100,LogLoss 随计算量增大稳定下降(约每 4× 计算量相对 LogLoss 改善约 0.1%),验证推荐域 Scaling Law。
四、PyTorch 简化实现
⚠️ 此为教学用简化版,完整工业实现需加入 Embedding 分片、混合精度、Rowwise Adagrad 等。
python
import torch
import torch.nn as nn
class FMB(nn.Module):
"""Factorization Machine Block with low-rank optimization."""
def __init__(self, n_in, n_out, d, k):
super().__init__()
self.k = k
# 低秩投影矩阵 Y: n_in -> k
self.Y = nn.Parameter(torch.randn(n_in, k))
self.ln = nn.LayerNorm(k)
self.mlp = nn.Sequential(
nn.Linear(k, n_out * d),
nn.ReLU()
)
self.n_out = n_out
self.d = d
def forward(self, X):
# X: [batch, n_in, d]
# 低秩 FM: X(X^T Y) -> [batch, n_in, k]
# 用 einsum 实现: batch i, n_in j, d k -> i,j,k
XY = torch.einsum('bnd,nk->bnk', X, self.Y) # [B, n_in, k]
fm = torch.einsum('bnd,bnk->bnk', X, XY) # [B, n_in, k]
fm = self.ln(fm)
flat = fm.reshape(fm.size(0), -1) # [B, n_in*k]
out = self.mlp(flat) # [B, n_out*d]
return out.reshape(fm.size(0), self.n_out, self.d)
class LCB(nn.Module):
"""Linear Compression Block."""
def __init__(self, n_in, n_out, d):
super().__init__()
self.proj = nn.Linear(n_in, n_out, bias=False)
self.d = d
def forward(self, X):
# X: [B, n_in, d] -> [B, n_out, d]
return self.proj(X.transpose(1, 2)).transpose(1, 2)
class WukongLayer(nn.Module):
def __init__(self, n_in, n_F, n_L, d, k):
super().__init__()
self.fmb = FMB(n_in, n_F, d, k)
self.lcb = LCB(n_in, n_L, d)
self.ln = nn.LayerNorm((n_F + n_L) * d)
def forward(self, X):
fmb_out = self.fmb(X) # [B, n_F, d]
lcb_out = self.lcb(X) # [B, n_L, d]
cat = torch.cat([fmb_out, lcb_out], dim=1) # [B, n_F+n_L, d]
cat_flat = cat.reshape(cat.size(0), -1)
# 残差连接(输入展平后与输出相加)
X_flat = X.reshape(X.size(0), -1)
if X_flat.shape == cat_flat.shape:
cat_flat = cat_flat + X_flat
return self.ln(cat_flat).reshape(cat.size(0), -1, self.d)
class Wukong(nn.Module):
def __init__(self, n_features, d=64, l=3, n_F=128, n_L=64, k=32, final_mlp=(256, 128)):
super().__init__()
self.embed = nn.Embedding(n_features, d)
self.stack = nn.Sequential(
*[WukongLayer(n_F + n_L, n_F, n_L, d, k) for _ in range(l)]
)
self.pool = nn.AdaptiveAvgPool1d(1)
self.final_mlp = nn.Sequential(
nn.Linear(d, final_mlp[0]),
nn.ReLU(),
nn.Linear(final_mlp[0], final_mlp[1]),
nn.ReLU(),
nn.Linear(final_mlp[1], 1),
nn.Sigmoid()
)
def forward(self, idx):
# idx: [B, n_features]
X = self.embed(idx) # [B, n_features, d]
# 第一层输入宽度需与 WukongLayer 期望的 n_in 对齐(此处简化:先投影到 n_F+n_L)
# 实际论文中第一层 n_in = n_features,后续层 n_in = n_F+n_L
Z = self.stack(X) # [B, n_F+n_L, d]
Z = self.pool(Z.permute(0, 2, 1)).squeeze(-1) # [B, d]
return self.final_mlp(Z)
说明:简化版中第一层
n_in = n_features,后续层n_in = n_F + n_L;完整复现需按论文设置各层n_in并初始化第一层投影。
五、关键点回顾
| 方面 | 要点 |
|---|---|
| 数学核心 | FM 低秩优化 O(n²d)→O(nkd),堆叠层使交互阶数 2ˡ 增长 |
| 结构创新 | FMB(二阶交叉)+ LCB(线性保阶)+ 残差 LN,可任意加深 |
| Scaling Law | Dense Scaling(扩 l, n_F, k, MLP)在 2 个数量级 GFLOP 上 Loss 稳定下降 |
| 意义 | 推荐模型从"Sparse Scaling(扩 Embedding)"转向"Dense Scaling(扩交互计算)" |