Meta-WHALE:逐层融合统一推荐模型

  • 论文标题: WHALE: A Scalable Unified Model for Recommendation with Wukong-HSTU Architecture

  • 论文链接: https://arxiv.org/pdf/2607.17017

  • 论文作者: Renqin Cai、Velvin Fu、Dawei Sun、Maggie Zhuang、Yuanjun Yao、Yu Shi、Zhiyong Wang、Zhongnan Fang、Xuan Cao、Jing Qian、Rui Li(Meta Platforms, Inc.)

  • 一句话总结: 把"非序列特征高阶交互(Wukong)"与"超长行为序列建模(HSTU)"做成逐层并行+逐层注意力交互的统一骨干,让每一层的高阶交互都能按候选/上下文去长序列里"检索证据",并通过系统侧 Triton/编译优化把它推到可线上部署。

背景与动机

推荐排序信号大体可以分成两类:一类是非序列特征(用户画像、候选内容属性、上下文、显式 user-item cross 等)的组合交互;另一类是用户行为历史序列(超长、强时序、带动作类型/作者/时间戳等 side info)。近年来两条路线各自把"规模"做到了很高:Wukong 代表了可扩展的高阶特征交互骨干,HSTU 代表了可扩展的超长行为序列骨干。

但实际做排序时,这两类信号不是"简单拼接就行"。很多请求级相关性,往往需要"某个候选+某个上下文"去历史里找特定的、细粒度的证据:可能是刚刚发生的一个动作、也可能是某个作者相关的重复偏好、也可能是一个跨时间段的行为模式。论文认为浅层混合(先把序列压成少量 summary embedding,再喂给非序列交互网络)会损失"按候选/上下文动态检索细粒度事件"的能力,因为一旦被压缩成 summary,后续的高阶交互就很难再精准回到"哪一条历史事件"上。

WHALE 的核心动机是:让 Wukong 和 HSTU 不只是"串一下",而是每一层都并行存在,并在每一层发生一次可学习的交互。这样,高阶交互可以反复从长序列里取证据,长序列表示也能随层数更深而更好地承载"可被检索"的信息。

整体架构

端到端数据流可以按"两个分支 + 每层融合"来理解:

  • 输入层把原始特征映射成两组 token:非序列 token(用户、候选、上下文、cross 等)与序列 token(历史行为序列,每个行为含 item+side info)。

  • 每一层同时做三件事:

    1. Wukong 分支在非序列 token 上做高阶交互,得到"交互后的非序列表示";

    2. HSTU 分支在序列 token 上做长序列建模,得到"上下文化的每步行为表示";

    3. Fusion 分支用交互表示当 Query,对长序列表示做注意力检索,取回与每个交互 token 最相关的历史证据,并把证据注入回非序列表示。

  • 下一层继续:非序列走"融合后的表示",序列走"HSTU 的输出",如此堆叠形成"逐层交换"。

下面进入逐模块拆解。为了讲清楚张量 Shape,我先统一符号。

符号约定(贯穿全文):

  • B:batch size

  • N:非序列特征 token 数(输入层融合后得到的非序列字段数)

  • L:行为序列长度(可到 15k)

  • D:embedding 维度(论文在实验里用过 64/128/256/512 等)

  • 非序列表示:X_ns^(ℓ) ∈ R^{B×N×D}

  • 序列表示:X_s^(ℓ) ∈ R^{B×L×D}

逐模块方案拆解(公式 + 变量含义 + Shape)

3.1 输入层:把异构特征映射成两组 token

模块作用:把原始 categorical / numerical / sequence 特征统一变成可被 Transformer/FM 处理的 dense embedding。

输入:

  • 非序列特征:用户/候选/上下文/显式 cross 等(包含 categorical 与 numerical)

  • 序列特征:用户历史行为序列,每个行为包含 item id + side info(动作类型、作者 id、时间戳等)

输出:

  • 非序列 token:E_ns ∈ R^{B×N×D}

  • 序列 token:E_s ∈ R^{B×L×D}

计算方式(按论文描述):

  • categorical 字段用 embedding lookup 表;

  • numerical 字段用一个 MLP tokenizer 映射到 embedding;

  • 序列侧每个行为把 item embedding 与 side-info embedding 通过 MLP 融合成一个 D 维向量。

变量说明:

  • N 是输入层完成融合后的非序列字段数;它不是"用户字段数+候选字段数"的简单加和,工程里可以把 cross/上下文也当成 token。

  • L 是行为历史长度,论文实验里做到了 3k/6k/10k/15k。

初始化:

Xns(0)=Ens,Xs(0)=Es X_{ns}^{(0)} = E_{ns},\quad X_{s}^{(0)} = E_sXns(0)=Ens,Xs(0)=Es

3.2 Wukong 分支:在非序列 token 上做高阶交互

模块作用:对非序列特征做可扩展的高阶交互建模,生成"候选+上下文相关"的交互表征,这些表征在 WHALE 里还会充当"去序列里检索证据"的 Query。

输入 / 输出:

  • 输入:X_ns^(ℓ) ∈ R^{B×N×D}

  • 输出:H_w^(ℓ) ∈ R^{B×N×D}

论文这里把 Wukong 作为黑盒骨干引用(在 Related Work 给了它的层级表达),在 WHALE 的单层里可以抽象成:

Hw(ℓ)=Wukong_module(Xns(ℓ))∈RB×N×D H_w^{(\ell)} = \mathrm{Wukong\module}(X{ns}^{(\ell)})\in \mathbb{R}^{B\times N\times D}Hw(ℓ)=Wukong_module(Xns(ℓ))∈RB×N×D

变量说明:

  • H_w^(ℓ) 的每个 token 可以理解为"某个字段/某个 cross 字段"经过多层交互后得到的高阶交互表示。

  • 由于 Wukong 的递归堆叠,增大深度/宽度会增加"交互阶数"和可表达性。

3.3 HSTU 分支:在超长行为序列上做序列建模

模块作用:把长度 L 可达 15k 的历史行为序列,编码成每个时间步都有上下文的表示,保留足够的细粒度证据,供后续 Fusion 检索。

输入 / 输出:

  • 输入:X_s^(ℓ) ∈ R^{B×L×D}

  • 输出:H_s^(ℓ) ∈ R^{B×L×D}

论文同样把 HSTU 作为可扩展序列骨干引用,在 WHALE 的单层里抽象成:

Hs(ℓ)=HSTU_module(Xs(ℓ))∈RB×L×D H_s^{(\ell)} = \mathrm{HSTU\_module}(X_s^{(\ell)})\in \mathbb{R}^{B\times L\times D}Hs(ℓ)=HSTU_module(Xs(ℓ))∈RB×L×D

变量说明:

  • H_s^(ℓ) 仍然是 L 步的表示(而不是被压成 1~k 个 summary),这是 WHALE 能"按交互 token 检索细粒度事件"的前提。

3.4 Fusion 分支:用非序列交互表示 Query 序列表示(核心创新点)

模块作用:让每个"高阶交互 token"都能从长行为序列里检索与自身最相关的证据,并把证据融合回非序列侧。

3.4.1 Cross-Attention(Wukong 作 Query,HSTU 作 Key/Value)

输入:

  • Wukong 输出(Query 源):H_w^(ℓ) ∈ R^{B×N×D}

  • HSTU 输出(Key/Value 源):H_s^(ℓ) ∈ R^{B×L×D}

预归一化 + 线性投影:

H~w(ℓ)=LNw(ℓ)(Hw(ℓ)),H~s(ℓ)=LNs(ℓ)(Hs(ℓ)) \tilde H_w^{(\ell)} = \mathrm{LN}_w^{(\ell)}(H_w^{(\ell)}),\quad \tilde H_s^{(\ell)} = \mathrm{LN}_s^{(\ell)}(H_s^{(\ell)})H~w(ℓ)=LNw(ℓ)(Hw(ℓ)),H~s(ℓ)=LNs(ℓ)(Hs(ℓ))

Q(ℓ)=H~w(ℓ)WQ(ℓ),K(ℓ)=H~s(ℓ)WK(ℓ),V(ℓ)=H~s(ℓ)WV(ℓ) Q^{(\ell)} = \tilde H_w^{(\ell)} W_Q^{(\ell)},\quad K^{(\ell)} = \tilde H_s^{(\ell)} W_K^{(\ell)},\quad V^{(\ell)} = \tilde H_s^{(\ell)} W_V^{(\ell)}Q(ℓ)=H~w(ℓ)WQ(ℓ),K(ℓ)=H~s(ℓ)WK(ℓ),V(ℓ)=H~s(ℓ)WV(ℓ)

张量 Shape:

  • Q^(ℓ) ∈ R^{B×N×D}

  • K^(ℓ), V^(ℓ) ∈ R^{B×L×D}

  • W_Q^(ℓ), W_K^(ℓ), W_V^(ℓ) ∈ R^{D×D}

注意力输出(把长序列证据汇聚到每个非序列交互 token 上):

A(ℓ)=softmax(Q(ℓ)(K(ℓ))⊤D) V(ℓ)∈RB×N×D A^{(\ell)} = \mathrm{softmax}\left(\frac{Q^{(\ell)} (K^{(\ell)})^\top}{\sqrt{D}}\right)\, V^{(\ell)}\in \mathbb{R}^{B\times N\times D}A(ℓ)=softmax(D Q(ℓ)(K(ℓ))⊤)V(ℓ)∈RB×N×D

变量说明:

  • Q(K)^T 的 shape 是 R^{B×N×L},表示每个非序列 token 对历史 L 个行为的注意力分布;

  • A^(ℓ) 的每个 token 可以理解为"该交互 token 从历史里取回的证据向量"。

3.4.2 Fusion MLP:把"交互表示 + 检索证据"合成新的非序列表示

论文把 Fusion MLP 写成"拼接后投影 + 残差"的形式:

Aˉ(ℓ)=Hw(ℓ)+Hw(ℓ) ∥ A(ℓ)WF(ℓ)∈RB×N×D \bar A^{(\ell)} = H_w^{(\ell)} + H_w\^{(\\ell)}\\,\\\|\\,A\^{(\\ell)} W_F^{(\ell)}\in \mathbb{R}^{B\times N\times D}Aˉ(ℓ)=Hw(ℓ)+Hw(ℓ)∥A(ℓ)WF(ℓ)∈RB×N×D

变量说明:

  • [\cdot\|\cdot] 是在最后一维拼接,shape 从 D 变成 2D

  • W_F^(ℓ) ∈ R^{2D×D} 把拼接后的表示投影回 D

  • 残差项 H_w^(ℓ) 表示"保留原始交互信息,再注入检索到的证据"。

3.4.3 SwiGLU FFN:对融合表示做非线性变换(带门控)

论文在 fusion 后接一个 pre-norm 的 SwiGLU FFN:

A^(ℓ)=LNffn(ℓ)(Aˉ(ℓ)) \hat A^{(\ell)} = \mathrm{LN}_{ffn}^{(\ell)}(\bar A^{(\ell)})A^(ℓ)=LNffn(ℓ)(Aˉ(ℓ))

Hf(ℓ)=Aˉ(ℓ)+SwiGLU_FFN(ℓ)(A^(ℓ))∈RB×N×D H_f^{(\ell)} = \bar A^{(\ell)} + \mathrm{SwiGLU\_FFN}^{(\ell)}(\hat A^{(\ell)})\in \mathbb{R}^{B\times N\times D}Hf(ℓ)=Aˉ(ℓ)+SwiGLU_FFN(ℓ)(A^(ℓ))∈RB×N×D

其中 SwiGLU 的一种写法(论文给的是 gate/up/down 三个投影):

SwiGLU_FFN(x)=SiLU(xWgate)⊙(xWup) Wdown \mathrm{SwiGLU\FFN}(x) = \mathrm{SiLU}(x W{gate}) \odot (x W_{up})\, W_{down}SwiGLU_FFN(x)=SiLU(xWgate)⊙(xWup)Wdown

变量说明:

  • SiLU(·):激活函数;

  • \odot:逐元素乘;

  • W_gate, W_up, W_down:可学习投影矩阵(shape 以实现为准,直觉上是把 D 维映射到更宽的中间维再回到 D)。

3.4.4 下一层的状态传递(两条链路都保留)

WHALE 的"递归层"把两条分支都持续保留:

Xns(ℓ+1)=Hf(ℓ),Xs(ℓ+1)=Hs(ℓ) X_{ns}^{(\ell+1)} = H_f^{(\ell)},\quad X_s^{(\ell+1)} = H_s^{(\ell)}Xns(ℓ+1)=Hf(ℓ),Xs(ℓ+1)=Hs(ℓ)

这点很关键:序列分支不会被压缩掉,因此下一层依然能在 L 级别的细粒度上被检索;非序列分支则带着"已注入的证据"进入下一层做更高阶的交互。

训练目标 / 指标口径

论文的离线主任务是平台主 engagement 预估任务,记二分类标签 y_i ∈ {0,1},模型输出概率 p_i。训练目标通常是 log loss(交叉熵):

Lce=−1B∑i=1B(yilog⁡pi+(1−yi)log⁡(1−pi)) \mathcal{L}{ce} = -\frac{1}{B}\sum{i=1}^{B} \left(y_i \log p_i + (1-y_i)\log(1-p_i)\right)Lce=−B1i=1∑B(yilogpi+(1−yi)log(1−pi))

论文汇报的是归一化熵(NE),用全量评估集的经验点击率 \bar p 做归一化:

NE=−1N∑i=1N(yilog⁡pi+(1−yi)log⁡(1−pi))−(pˉlog⁡pˉ+(1−pˉ)log⁡(1−pˉ)) NE = \frac{-\frac{1}{N}\sum_{i=1}^{N}\left(y_i \log p_i + (1-y_i)\log(1-p_i)\right)}{-\left(\bar p \log \bar p + (1-\bar p)\log(1-\bar p)\right)}NE=−(pˉlogpˉ+(1−pˉ)log(1−pˉ))−N1∑i=1N(yilogpi+(1−yi)log(1−pi))

变量说明:

  • 分子是普通 log loss 的期望;

  • 分母是"只预测常数 CTR=\bar p"的熵(相当于把难度做了归一化);

  • NE 越低越好 。论文提到在该平台上 0.05% 的 NE 增益就算明显提升

系统与效率优化(把可训练/可上线做实)

WHALE 在结构上引入了"非序列 Query × 超长序列 KV"的 cross-attention,这在工程上很容易变成吞吐瓶颈。论文把一大段篇幅用于说明:怎么把它从"纸上可行"变成"线上可跑"。

6.1 训练侧:面向 QPS 的优化

6.1.1 Triton 自定义 cross-attention kernel

场景特征:N 相对小(例如 64 个非序列 token),L 很大(例如 15,000 的历史)。这种 Q 小、KV 巨长 的 attention 很容易内存带宽受限。

论文给了两类关键优化:

  • 共享 Key/Value(K=V):把 value 直接复用 key,减少一次大张量读写与反传的内存流量。注意力形态仍是

A(ℓ)=softmax(Q(ℓ)(K(ℓ))⊤D)K(ℓ) A^{(\ell)} = \mathrm{softmax}\left(\frac{Q^{(\ell)} (K^{(\ell)})^\top}{\sqrt{D}}\right)K^{(\ell)}A(ℓ)=softmax(D Q(ℓ)(K(ℓ))⊤)K(ℓ)

  • 不对称反传调度(KV-parallel vs Q-parallel) :因为 N \ll L,反传分块策略会显著影响归约与带宽开销。论文描述基于形状在运行时选择两种 backward schedule,据称带来约 1.2x kernel 级加速。
6.1.2 Shared-gate SwiGLU:把 3 次 GEMM 减到 2 次

标准 SwiGLU 需要 gate/up/down 三个投影矩阵,带来 3 次矩阵乘。论文使用 共享 gate 与 up 的变体(把两者权重绑定为同一个 W_s),从 3 次 GEMM 降到 2 次,理论上减少约 33% 的 FFN 乘法开销。

SharedGateSwiGLU(x)=(SiLU(xWs)⊙(xWs))Wdown \mathrm{SharedGateSwiGLU}(x) = \left(\mathrm{SiLU}(xW_s)\odot (xW_s)\right) W_{down}SharedGateSwiGLU(x)=(SiLU(xWs)⊙(xWs))Wdown

6.1.3 混合精度与编译优化
  • dense backbone(Wukong/HSTU/Fusion/FFN)用 BF16;

  • 对数值敏感部分(loss、task head 等)保留 FP32;

  • 对 dense 子图使用 torch.compile / TorchInductor 做算子融合与内存规划。

论文声称训练侧优化总体带来约 30% 的训练吞吐提升(相对未优化版本)。

6.2 推理侧:面向 QPS/延迟的优化

推理的额外难点是动态 shape、kernel launch overhead、以及 CPU-GPU 同步(shape 计算导致)。论文提到:

  • shape-aware kernel tuning + AOTInductor 做更激进的融合,报告 18% 更高吞吐

  • 用 shape hint tensors 降低同步点,单次同步消除可降 1--5ms,最终报告 推理 QPS +15%(相对未优化版本)。

实验与分析

7.1 离线设置

数据与训练:

  • 来自某大型短视频社交平台的训练日志;

  • 训练样本约 80B,评估样本约 4B;

  • 每个模型训练 1 个 epoch;

  • 多数实验在 NVIDIA B200 上完成。

7.2 主结果:WHALE vs Wukong-only / HSTU-only

解读要点(紧贴图 3):

  • 论文对齐 FLOPs(通过调 hidden dim、层数、序列长度等),在 8 / 14 / 32 GFLOPs 三个点上比较;

  • WHALE 在同等复杂度下相对更稳地拿到更好的 NE 增益,说明"统一建模"不是靠堆某一条分支,而是让两条分支互补。

7.3 Scaling:更长序列、更深层数、更宽维度都继续涨

这里最关键的不是某个绝对数字,而是趋势:

  • 序列长度从 3k 提到 15k 仍然持续带来 NE 增益,说明架构确实能吃下超长历史;

  • 深度从 2 层加到 8 层持续增益,符合"逐层交换越多、交互越深"的直觉;

  • 宽度从 64 到 512 仍然继续涨,说明两条分支的容量扩张都能转化为收益。

7.4 Fusion 消融:注意力与逐层融合是硬骨头

论文把不同融合策略的质量回退(相对 WHALE)列成了表 1。这里我按原表重写成可读的结构化表格:

层级 融合策略 NE 回退(%)
Reference WHALE(逐层并行 + 逐层 attention 融合) 0.00
Architecture level Shallow-hybrid fusion(先压缩序列,再喂给交互网络的浅层融合) 0.25
Module level Avg.-pooling fusion(用序列平均池化替代注意力检索) 0.23
Module level Fusion w/o MLP projection(去掉 Fusion MLP 投影/残差) 0.11
Module level Fusion w/o SwiGLU FFN(去掉融合后的 FFN 非线性) 0.08

读表结论(从"差多少"反推"哪个更关键"):

  • "浅层融合"回退 0.25%,说明逐层交换不是装饰,而是有效的容量利用方式;

  • "Avg pool 替代注意力"回退 0.23%,说明检索式注意力是把长序列价值兑现的关键;

  • MLP 和 FFN 的作用也不小,说明"取回证据后如何融合/非线性加工"同样重要。

7.5 线上 A/B:有收益,但吞吐有代价

论文做了 14 天线上 A/B,结果见表 2(我按原表重写):

指标 相对变化
Primary evaluation metric +0.113%
Metric 1 +0.824%
Metric 2 +1.820%
Inference QPS -5%

这组结果本质上是"质量-成本"交换的一个落地例子:质量指标是正向提升,同时推理吞吐回退 5%。论文的表述是该回退仍在可接受的 serving budget 内,并且已经部署到生产系统。

优势与局限(只基于论文内容)

优势:

  • 结构上把两条可扩展骨干(Wukong / HSTU)做成可复用的统一层,天然支持三种 scaling knob:更长 L、更深层数、更宽 D

  • Fusion 机制是"检索式"的:每个交互 token 都能对长序列做 attention,而不是依赖固定 summary,因此更适合候选/上下文依赖的证据检索。

  • 工程上给了较完整的系统优化路径(自定义 kernel、共享 KV、shape-aware backward、shared-gate SwiGLU、编译优化),明确指向"可线上部署"。

局限与代价:

  • 线上实验显示推理 QPS 有回退(-5%),说明统一架构在 serving 侧仍有明显成本,需要依赖系统优化才能上量。

  • 文中很多结论建立在大规模工业数据与特定 serving stack 上;迁移到别的平台/任务时,最关键的仍然是:N\ll L 的 cross-attention 形态是否成立,以及能否复现相同级别的 kernel/编译优化收益。

相关推荐
墨舟的AI笔记1 小时前
微前端隔离边界:qiankun 与 Module Federation 的取舍之道
人工智能
字节逆旅1 小时前
有哪些工作不可以交给AI
人工智能·程序员
一路向北North1 小时前
Spring AI(6) :对话机器人-会话历史
java·人工智能·spring
2zcode1 小时前
基于MATLAB神经网络的心力衰竭预测与临床辅助决策系统研究
人工智能·神经网络·matlab
ACP广源盛139246256731 小时前
国产算力互联IX8024@ACP#Kimi K3 开源后的端侧部署硬件架构分析
大数据·人工智能·分布式·单片机·嵌入式硬件
江边风声2 小时前
从薄板到厚板、从吸盘到夹板边——坤鹏伯爵的取放技术体系是怎样覆盖全制程的
人工智能·科技·自动化·制造·pcb工艺
商业模式源码开发2 小时前
小米新车未发布即遭 AI 谣言攻击:黑色 GEO 的运作原理与企业正规防御方案
大数据·人工智能·ai·geo
edtoplort2 小时前
A股最大IPO长鑫科技295亿:从年亏163亿到日赚3亿
大数据·人工智能·科技
道影子2 小时前
《道德经》031兵者不祥,胜以丧礼处之
人工智能·深度学习·算法