Wukong: Towards a Scaling Law for Large-Scale Recommendation

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 数)和维度来增加参数量。 这种扩展方式存在两大核心问题:

  1. 无法增强高阶特征交互能力:单纯扩大稀疏嵌入参数并不能提升模型捕捉特征间复杂高阶交互的能力,随着嵌入表增大,收益迅速饱和,甚至出现边际效用递减。
  2. 硬件利用率低下: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(扩交互计算)"
相关推荐
硅基流动1 小时前
山东铁路基金公司与硅基流动达成战略合作,共建 Token 工厂
人工智能·科技
飞哥数智坊1 小时前
难道 AI 真要让程序员三班倒了?
人工智能·ai编程
玫瑰互动GEO1 小时前
抖音SEO优化技术拆解:搜索排名四大因子与4步落地算法分析
人工智能·算法·搜索引擎·语音识别
IT_陈寒1 小时前
Vite打包时踩了个坑,static资源去哪了?
前端·人工智能·后端
龙虾PRO1 小时前
2026 年 AI 智能体工具调用:ReAct 模式与函数调用怎么选才不踩坑
前端·人工智能·react.js
其实防守也摸鱼1 小时前
如何使用自动化教育SRC漏洞挖掘系统--AutoHunter
运维·网络·人工智能·学习·安全·web安全·自动化
微石科技1 小时前
体征监测设备厂家怎么选?宁波市微石科技:专业级硬件,全品类可定制
大数据·人工智能·科技
泡干脆面就番茄1 小时前
机器学习:TF-IDF
人工智能·机器学习·tf-idf
iori97king1 小时前
织信开发日志 18:从 informat-skills 看织信如何把平台能力交给 AI Agent
人工智能·低代码·织信