MARM

0 阅读4分钟

标题:MARM:通过记忆增强与可扩展复杂度解锁推荐Cache Scaling-Law
单位:快手
链接:MARM: Unlocking the Recommendation Cache Scaling-Law through Memory Augmentation and Scalable Complexity(2024.11)

摘要

过去几年,缩放定律一直指导着 GPT 等语言模型的设计,使人们能够根据可学习参数规模和训练样本规模估计模型的预期性能。值得注意的是,自然语言处理领域的缩放定律不能直接应用于推荐系统,原因如下:

  1. 训练样本量和模型参数量通常不是模型的瓶颈。我们的推荐系统每天能够生成超过 500 亿条用户样本,如此庞大的训练数据可以轻松支持模型参数量超过 2000 亿,超越许多参数量约为 1000 亿的大语言模型。
  2. 推荐系统必须谨慎控制浮点运算次数(FLOPs)。在训练阶段,我们每天需要处理海量推荐样本;在在线推理阶段,则必须在毫秒内响应,而大语言模型通常需要数秒。

根据推荐系统与大语言模型的上述差异可以得出结论:对于推荐系统模型,与参数量相比,FLOPs 是成本更高、必须谨慎控制的因素。

本文提出具有里程碑意义的工作 MARM(Memory Augmented Recommendation Model,记忆增强推荐模型),成功探索出一种新的缓存缩放定律。MARM 缓存复杂模块的部分计算结果,仅增加少量推理 FLOPs,便将基于单层注意力的序列兴趣建模模块扩展到多层,即将模块的时间复杂度从 O(n2d)O(n^2d) 降低为 O(nd)O(nd)。借助缓存思想,MARM 方案显著突破了计算瓶颈,可以无缝增强所有面向用户序列的兴趣提取模块,甚至还能增强其他模型。

为了支撑 MARM,我们构建了一个容量为 60 TB 的缓存存储中心,用于离线训练和在线服务。全面的实验结果表明,MARM 在线下使 GAUC 提升 0.43%0.43\%,在线上使用户人均播放时长提升 2.079%2.079\%。MARM 已部署到真实短视频平台,每天为数千万用户提供服务。

CCS 概念: 信息系统;推荐系统。

1 引言

近年来,模型规模和数据规模的扩展已成为自然语言处理、计算机视觉和推荐系统等多个领域的核心议题。根据 OpenAI 的缩放定律技术报告,随着数据量以及模型宽度和深度的增加,模型性能会按照特定幂律得到改善。在缩放定律理念的推动下,人们提出了许多基于 Transformer 的大模型,并取得显著效果,例如用于对话的 ChatGPT、用于多模态理解的 Qwen-VL,以及用于代码生成的 DeepSeek-Coder 等。

推荐系统领域也出现了验证缩放定律的尝试,例如 Wukong 和 HSTU,但这些方法对特征工程或模型架构作出了很强的假设,例如移除所有静态特征,或者用 Transformer 替换整个模型模块。事实上,推荐系统排序模型的常见设计范式始终是:首先精心构造数百个特征作为模型输入,再使用多任务预测模块,为用户与候选物品对预测点击、点赞、评论等多种分数,如图 1(a)所示。然而,近期的推荐系统缩放定律研究正在大幅改变这一学习范式,使其很难直接部署到真实的在线排序模型中。本文重点探索上述常用排序模型架构下的缩放定律。

image.png 图 1:(a)排序模型的简化示例。(b)流式数据训练中的性能差距。以在线基线模型进行热启动后,即便只对样本下采样 50%50\%,模型性能仍会持续下降。(c)参数分布。大语言模型的参数(蓝色)主要位于“稠密激活”的、基于 Transformer 的 DNN 模块中;排序模型的参数(绿色)则主要由“稀疏激活”的特征参数构成。

由于排序模型的架构和使用方式与纯 Transformer 大语言模型不同,我们总结出自然语言处理缩放定律不适用于推荐排序模型的三个原因。

训练样本

自然语言处理模型通常拥有稳定的数据语料库,可用于从头训练大语言模型,而工业推荐系统模型往往采用流式训练范式。在真实推荐系统中,我们的模型每天使用超过 500 亿条用户日志进行流式训练。

这里给出三个模型变体的训练情况:“在线基线模型”“使用 50%50\% 训练数据下采样的在线基线模型”和“冷启动模型”,其中在线基线模型已经训练了数年。从图 1(b)可以得到两个结论:第一,从头训练会显著损害性能;第二,即使模型参数由一个已经训练数年的模型进行热启动,对实时训练样本进行下采样仍会导致性能下降。这些现象证明,推荐系统模型对无限数据有强烈需求,只有这样才能捕捉用户的实时偏好。因此,在工业场景中分析训练数据规模并无必要。

可学习参数

如图 1(c)所示,大语言模型的参数主要位于“稠密激活”的、基于 Transformer 的 DNN 模块中,参数量约为 1000 亿,而“稀疏激活”的词元参数很少,不到 1000 万。排序模型的参数分布恰好相反:大部分可学习参数集中在“稀疏激活”的特征参数中,超过 2000 亿;“稠密激活”的 DNN 参数则少得多,约为 1 亿。

这里的“稀疏”是指每次只有一小部分参数被激活并参与计算,例如,在大语言模型中,一句话只需查询少量词元;在推荐系统中,一条用户与物品训练样本只需查询一个用户 ID 和一个物品 ID。“稠密”则指计算流程中的参数会被全部激活,用于估计最终结果。事实上,仅从可学习参数量看,我们的排序模型超过 2000 亿参数,远多于许多约有 1000 亿参数的大语言模型。这表明,可学习参数量并不是排序模型的瓶颈。

计算复杂度与 FLOPs

事实上,推荐系统模型的推理 FLOPs 远低于大语言模型,并且存在上限。这是因为模型需要高效处理数亿次请求,所以稠密激活 DNN 模块的规模要小得多,约为 1 亿参数。为了保证服务在毫秒内处理每次请求时的稳定性和鲁棒性,而大语言模型通常需要数秒,我们不能盲目增加推理复杂度。

此外,推荐模型训练的实时性会极大影响在线性能。在离线训练资源有限时,模型的计算复杂度会影响流式训练的实时性,从而与在线性能形成权衡。例如,如果离线模型产生 20 分钟延迟,在线性能便会下降 1%1\%

综合上述三个方面,与大语言模型相比,排序模型具有以下优势与劣势:拥有“无限”的流式数据、海量参数存储空间,以及 FLOPs 相对较低的 DNN 模块。换言之,对于排序模型,数据和存储资源相对便宜,计算资源却非常昂贵。受此启发,我们思考能否利用排序模型在数据和存储方面的优势,弥补其计算方面的劣势。也就是说,能否缓存复杂模块的部分计算结果,从而降低其时间复杂度?

为回答这一问题,我们提出具有里程碑意义的工作 MARM,它在缓存规模与模型性能之间实现了一种新的推荐系统缩放定律。MARM 扩展了工业排序模型中最重要的用户兴趣提取模块之一,该模块用于结合候选物品,计算用户历史物品的重要性,如图 1(a)底部所示。

据我们所知,在实现该模块时,DIN、SIM、SDIM 和 TWIN 等许多优秀方法都使用图 2(a)所示的单层目标注意力(Target Attention,TA)机制。直观来看,可以在最终 TA 层之前堆叠多层自注意力(Self-Attention,SA),如图 2(b)所示,以增强该模块。遗憾的是,SA 的时间复杂度为 O(n2d)O(n^2d),远高于 TA 的 O(nd)O(nd),其中 nn 表示序列长度,dd 表示表征维度。这种差异对长序列尤其敏感,例如 n>1000n>1000 时,会使 FLOPs 大幅增加。

在在线服务中,排序模型需要同时预测 50 至 8000 个用户与候选物品对,从而为一位用户选出最优的几十个物品。这会使朴素多层注意力机制产生大量重复计算,给系统带来沉重压力,如图 2(c)所示。

幸运的是,工业推荐系统广泛使用的流式数据具有时间顺序。若在最新样本上缓存频繁使用的掩码自注意力模块结果,就可以用简单 TA 层替代复杂的掩码 SA 层,从而高效预测不同候选物品,如图 2(d)所示。通过这种方式,MARM 只需少量 FLOPs,便能将基于单层注意力的序列兴趣建模模块扩展到多层。

具体而言,MARM 方案显著突破计算瓶颈,可以无缝增强所有面向用户序列的兴趣提取模块。我们还发现,MARM 缓存结果对检索模型、级联模型等其他模型也有帮助。在实验中,我们全面探索缓存规模与模型性能之间的缩放定律,并在缓存注意力深度、序列长度和嵌入维度之间寻找理想平衡。

image.png 图 2:MARM 的动机:利用缓存思想,将 FLOPs 较高的自注意力替换为 FLOPs 较低的target-attention。

本文的主要贡献如下:

  1. 首次从缓存与性能这一全新视角,在广泛使用的排序模型架构下探索新的缩放定律。
  2. 设计简单而高效的 MARM,利用排序模型在数据和存储方面的优势,降低掩码自注意力的高时间复杂度。我们成功地以很小的推理 FLOPs 开销,将基于单层注意力的序列兴趣建模模块扩展到多层。
  3. 证明 MARM 具有很强的适应性和可扩展性,可以无缝集成到现有高性能 Transformer 模型中,并推动模型平滑过渡到 GPT 风格。我们还在线下和线上进行了广泛实验与对比消融研究,平均实现线下 GAUC 提升 0.43%0.43\%、线上用户人均播放时长提升 2.079%2.079\%

2 方法

本节不介绍排序模型完整架构或损失函数的所有细节。为了便于理解,我们只深入讨论 MARM 模块,并结合图 2(d)说明其工作方式。

2.1 MARM 工作流程

本节介绍 MARM 在最新曝光物品序列上的工作流程。该方法构建多层、仅解码器 Transformer 架构,以捕捉用户相对于特定目标物品的兴趣。MARM 主要由四部分组成:

  1. 序列生成器,用于生成已向用户曝光的物品。
  2. 外部缓存存储,用于查询计算结果。
  3. 多层目标注意力模块,用于计算结果。
  4. 将中间结果发送到缓存存储。

2.1.1 序列生成器

给定任意用户 ID UidX\mathrm{UidX},假设序列生成器能够按时间顺序生成其最新曝光物品序列:

[Iid1,,Iidn]=ExposureSeqGen(UidX,n).(1)[\mathrm{Iid1},\ldots,\mathrm{Iidn}]=\operatorname{ExposureSeqGen}(\mathrm{UidX},n).\tag{1}

其中,nn 表示序列长度,[Iid1,,Iidn][\mathrm{Iid1},\ldots,\mathrm{Iidn}] 是用户 UidX\mathrm{UidX} 的行为序列。由于该序列包含清晰的用户反馈,我们也可以只使用满足特定条件的物品,例如只使用带有长播标签的物品;这一扩展见第 2.2 节。

2.1.2 查询缓存存储

如图 2 所示,给定 UidX\mathrm{UidX} 的序列 [Iid1,,Iidn][\mathrm{Iid1},\ldots,\mathrm{Iidn}] 和候选物品 IidY\mathrm{IidY},MARM 需要两类表征:可学习的基于 ID 的嵌入,以及不可学习的缓存结果。

可学习嵌入。 在推荐系统中,基于 ID 的嵌入属于模型的稀疏激活特征,可以直接查询:

IidY=SparseFeatureLookUp(IidY),[Iid1,,Iidn]=SparseFeatureLookUp([Iid1,,Iidn]).(2)\begin{aligned} \mathbf{IidY}&=\operatorname{SparseFeatureLookUp}(\mathrm{IidY}),\\ [\mathbf{Iid1},\ldots,\mathbf{Iidn}]&=\operatorname{SparseFeatureLookUp}([\mathrm{Iid1},\ldots,\mathrm{Iidn}]). \end{aligned}\tag{2}

其中,IidYRF\mathbf{IidY}\in\mathbb{R}^{F}FF 是控制特征维度的超参数。[Iid1,,Iidn][\mathbf{Iid1},\ldots,\mathbf{Iidn}] 是图 2(a)底部目标注意力的输入。需要注意的是,已观看物品还可以包含作者 ID、标签和用户交互标签等其他属性。为简化说明,本文只保留物品 ID 来表示这些物品。

不可学习的缓存结果。 我们已经在外部键值存储中缓存了 LL 层的中间结果,因此需要准确查询这些结果。为此,我们设计了一种哈希策略,用于生成能够表达如下语义的“键”:用户与物品对在第 ii 层的缓存值。例如,给定用户 UidX\mathrm{UidX} 及其物品序列 [Iid1,][\mathrm{Iid1},\ldots],若要查找深度为 ii 的缓存结果,可以按下式生成哈希键列表:

[Iid1UidXi,]=UserItemDepthHash(UidX,[Iid1,],i).(3)[\mathrm{Iid1}_{\mathrm{UidX}}^{i},\ldots]=\operatorname{UserItemDepthHash}(\mathrm{UidX},[\mathrm{Iid1},\ldots],i).\tag{3}

其中,Iid1UidXiZ\mathrm{Iid1}_{\mathrm{UidX}}^{i}\in\mathbb{Z} 表示用于访问正确且唯一缓存结果的哈希键。基于这些键,可以得到:

[Iid1UidXi,]=MARMCacheLookUp([Iid1UidXi,]).(4)[\mathbf{Iid1}_{\mathrm{UidX}}^{i},\ldots]=\operatorname{MARMCacheLookUp}([\mathrm{Iid1}_{\mathrm{UidX}}^{i},\ldots]).\tag{4}

其中,Iid1UidXiRd\mathbf{Iid1}_{\mathrm{UidX}}^{i}\in\mathbb{R}^{d}dd 为注意力模块的维度。将所有层的缓存结果组合起来,可以得到:

[Iid1UidX1Iid2UidX1Iid3UidX1IidnUidX1Iid1UidX2Iid2UidX2Iid3UidX2IidnUidX2Iid1UidXLIid2UidXLIid3UidXLIidnUidXL].\begin{bmatrix} \mathbf{Iid1}_{\mathrm{UidX}}^{1} & \mathbf{Iid2}_{\mathrm{UidX}}^{1} & \mathbf{Iid3}_{\mathrm{UidX}}^{1} & \cdots & \mathbf{Iidn}_{\mathrm{UidX}}^{1}\\ \mathbf{Iid1}_{\mathrm{UidX}}^{2} & \mathbf{Iid2}_{\mathrm{UidX}}^{2} & \mathbf{Iid3}_{\mathrm{UidX}}^{2} & \cdots & \mathbf{Iidn}_{\mathrm{UidX}}^{2}\\ \vdots & \vdots & \vdots & \ddots & \vdots\\ \mathbf{Iid1}_{\mathrm{UidX}}^{L} & \mathbf{Iid2}_{\mathrm{UidX}}^{L} & \mathbf{Iid3}_{\mathrm{UidX}}^{L} & \cdots & \mathbf{Iidn}_{\mathrm{UidX}}^{L} \end{bmatrix}.

2.1.3 多层目标注意力

至此,我们已经获得可学习的目标注意力 ID 嵌入,以及后续不可学习的目标注意力缓存结果。首先根据可学习嵌入,通过 TA 机制结合特定目标物品信息 IidY\mathbf{IidY} 进行计算:

IidYUidX1=TargetAttention(IidY,[Iid1,]).(5)\mathbf{IidY}_{\mathrm{UidX}}^{1}=\operatorname{TargetAttention}(\mathbf{IidY},[\mathbf{Iid1},\ldots]).\tag{5}

随后,将结果 IidYUidX1Rd\mathbf{IidY}_{\mathrm{UidX}}^{1}\in\mathbb{R}^{d} 与不可学习的缓存结果一起输入后续层,通过 FLOPs 较低、复杂度为 O(nd)O(nd) 的 TA,模拟 FLOPs 较高、复杂度为 O(n2d)O(n^2d) 的掩码自注意力:

IidYUidX2=TargetAttention1(IidYUidX1,[Iid1UidX1,]),IidYUidX3=TargetAttention2(IidYUidX2,[Iid1UidX2,]), IidYUidXL+1=TargetAttentionL(IidYUidXL,[Iid1UidXL,]).(6)\begin{aligned} \mathbf{IidY}_{\mathrm{UidX}}^{2} &=\operatorname{TargetAttention}^1\left(\mathbf{IidY}_{\mathrm{UidX}}^{1},[\mathbf{Iid1}_{\mathrm{UidX}}^{1},\ldots]\right),\\ \mathbf{IidY}_{\mathrm{UidX}}^{3} &=\operatorname{TargetAttention}^2\left(\mathbf{IidY}_{\mathrm{UidX}}^{2},[\mathbf{Iid1}_{\mathrm{UidX}}^{2},\ldots]\right),\\ &\ \vdots\\ \mathbf{IidY}_{\mathrm{UidX}}^{L+1} &=\operatorname{TargetAttention}^L\left(\mathbf{IidY}_{\mathrm{UidX}}^{L},[\mathbf{Iid1}_{\mathrm{UidX}}^{L},\ldots]\right). \end{aligned}\tag{6}

其中,IidYUidXL+1Rd\mathbf{IidY}_{\mathrm{UidX}}^{L+1}\in\mathbb{R}^{d} 表示属于 (UidX,IidY)(\mathrm{UidX},\mathrm{IidY}) 对的最终序列兴趣建模输出,随后会被送入多任务学习模块,以预测真实标签。

2.1.4 将结果写入缓存

前文已经介绍 MARM 的计算过程,下面说明如何将中间结果保存到缓存存储。与之前类似,首先生成“哈希键”来标识结果的含义,例如 (UidX,IidY)(\mathrm{UidX},\mathrm{IidY}) 对在不同层 [1,,L][1,\ldots,L] 上的结果:

[IidYUidX1,,IidYUidXL]=UserItemDepthHash(UidX,IidY,[1,,L]).(7)[\mathrm{IidY}_{\mathrm{UidX}}^{1},\ldots,\mathrm{IidY}_{\mathrm{UidX}}^{L}]=\operatorname{UserItemDepthHash}(\mathrm{UidX},\mathrm{IidY},[1,\ldots,L]).\tag{7}

随后,按照顺序将值及其对应的键写入缓存:

MARMCacheSave([IidYUidX1,,IidYUidXL],[IidYUidX1,,IidYUidXL]).(8)\operatorname{MARMCacheSave}\left([\mathrm{IidY}_{\mathrm{UidX}}^{1},\ldots,\mathrm{IidY}_{\mathrm{UidX}}^{L}],[\mathbf{IidY}_{\mathrm{UidX}}^{1},\ldots,\mathbf{IidY}_{\mathrm{UidX}}^{L}]\right).\tag{8}

在流式训练中,随着实时计算结果被持续写入,MARM 缓存会逐步积累所有用户历史行为在 MARM 每一层上的计算结果。通过缓存的自然积累,MARM 能够以很小的推理 FLOPs 运行仅解码器 Transformer 架构。算法 1 给出了 MARM 的伪代码,图 3 还展示了流式场景中的用户观看示例。

image.png 图 3:MARM 实时缓存的更新与使用流程。

算法 1:MARM 训练与缓存积累过程

image.png

输入: 序列长度 nn、MARM 深度 LL、表征维度 dd
输出: 更新后的 MARM 缓存存储

  1. 对数据流中的每个 (UidX,IidY)(\mathrm{UidX},\mathrm{IidY}) 执行以下步骤。
  2. 生成序列:[Iid1,,Iidn]=ExposureSeqGen(UidX,n)[\mathrm{Iid1},\ldots,\mathrm{Iidn}]=\operatorname{ExposureSeqGen}(\mathrm{UidX},n)
  3. 准备 MARM 输入:
  4. IidY=ModelSparseFeatureLookUp(IidY)\mathbf{IidY}=\operatorname{ModelSparseFeatureLookUp}(\mathrm{IidY})
  5. [Iid1,,Iidn]=ModelSparseFeatureLookUp([Iid1,,Iidn])[\mathbf{Iid1},\ldots,\mathbf{Iidn}]=\operatorname{ModelSparseFeatureLookUp}([\mathrm{Iid1},\ldots,\mathrm{Iidn}])
  6. i=1i=1LL
  7. 生成键 [Iid1UidXi,,IidnUidXi]=UserItemDepthHash(UidX,[Iid1,],i)[\mathrm{Iid1}_{\mathrm{UidX}}^{i},\ldots,\mathrm{Iidn}_{\mathrm{UidX}}^{i}]=\operatorname{UserItemDepthHash}(\mathrm{UidX},[\mathrm{Iid1},\ldots],i)
  8. 查询缓存 [Iid1UidXi,,IidnUidXi]=MARMCacheLookUp([Iid1UidXi,])[\mathbf{Iid1}_{\mathrm{UidX}}^{i},\ldots,\mathbf{Iidn}_{\mathrm{UidX}}^{i}]=\operatorname{MARMCacheLookUp}([\mathrm{Iid1}_{\mathrm{UidX}}^{i},\ldots])
  9. 执行 MARM 前向计算:
  10. IidYUidX1=TargetAttention(IidY,[Iid1,,Iidn])\mathbf{IidY}_{\mathrm{UidX}}^{1}=\operatorname{TargetAttention}(\mathbf{IidY},[\mathbf{Iid1},\ldots,\mathbf{Iidn}])
  11. i=1i=1LL
  12. IidYUidXi+1=TargetAttentioni(IidYUidXi,[Iid1UidXi,,IidnUidXi])\mathbf{IidY}_{\mathrm{UidX}}^{i+1}=\operatorname{TargetAttention}_i(\mathbf{IidY}_{\mathrm{UidX}}^{i},[\mathbf{Iid1}_{\mathrm{UidX}}^{i},\ldots,\mathbf{Iidn}_{\mathrm{UidX}}^{i}])
  13. 将结果保存到缓存:
  14. 生成键 [IidYUidX1,,IidYUidXL]=UserItemDepthHash(UidX,IidY,[1,,L])[\mathrm{IidY}_{\mathrm{UidX}}^{1},\ldots,\mathrm{IidY}_{\mathrm{UidX}}^{L}]=\operatorname{UserItemDepthHash}(\mathrm{UidX},\mathrm{IidY},[1,\ldots,L])
  15. 执行 MARMCacheSave([IidYUidX1,,IidYUidXL],[IidYUidX1,,IidYUidXL])\operatorname{MARMCacheSave}([\mathrm{IidY}_{\mathrm{UidX}}^{1},\ldots,\mathrm{IidY}_{\mathrm{UidX}}^{L}],[\mathbf{IidY}_{\mathrm{UidX}}^{1},\ldots,\mathbf{IidY}_{\mathrm{UidX}}^{L}])
  16. 处理下一条流式数据。

2.2 MARM 与 SIM 结合

第 2.1 节讨论了朴素 MARM 设置,即在用户最新观看物品序列上建模其兴趣。然而,公式(1)中的序列长度 nn 有限,约为 100。一旦较早存储的结果超出这一范围,就不会再被使用,从而造成存储资源浪费。本节进一步将 MARM 与基于搜索的兴趣模型 SIM 结合,以最大限度发挥缓存结果的作用。

形式上,SIM 提出了一种两阶段级联建模范式:首先引入粗粒度的通用搜索单元(General Search Unit,GSU),针对目标物品回溯用户的终身历史,例如长度超过 10000,并搜索相关性最高的 Top-KK 物品序列;随后使用细粒度的精确搜索单元(Exact Search Unit,ESU)压缩搜索到的序列信息,得到用户相对于目标物品的精确兴趣。

过去几年中,SIM 一直是推动排序模型迭代的重要引擎。近期工作 TWIN 在两个阶段中使用共享的 GSU 和 ESU 模块,实现无偏建模。受 TWIN 启发,我们也考虑使用同步 GSU 模块,为每一层搜索相关性最高的 Top-KK 缓存结果,如图 4(a)所示。

给定用户 UidX\mathrm{UidX} 最新的长期序列 [Iid1,,Iid10000][\mathrm{Iid1},\ldots,\mathrm{Iid10000}],MARM GSU 的目标是生成多层搜索输入序列。例如:

[S_Iid1i,,S_IidKi]=MARMGSU(IidY,[Iid1,,Iid10000],K,i).(9)[\mathrm{S\_Iid1}^{i},\ldots,\mathrm{S\_IidK}^{i}]=\operatorname{MARMGSU}(\mathrm{IidY},[\mathrm{Iid1},\ldots,\mathrm{Iid10000}],K,i).\tag{9}

其中,MARMGSU\operatorname{MARMGSU} 与排序模型共享相同的目标注意力参数,[S_Iid1i,,S_IidKi][\mathrm{S\_Iid1}^{i},\ldots,\mathrm{S\_IidK}^{i}] 表示第 ii 层中对候选物品 IidY\mathrm{IidY} 具有最高注意力权重的 Top-KK 搜索物品。将所有层的搜索键组合起来,可以得到:

[S_Iid11S_Iid21S_Iid31S_IidK1S_Iid12S_Iid22S_Iid32S_IidK2S_Iid1LS_Iid2LS_Iid3LS_IidKL].\begin{bmatrix} \mathrm{S\_Iid1}^{1} & \mathrm{S\_Iid2}^{1} & \mathrm{S\_Iid3}^{1} & \cdots & \mathrm{S\_IidK}^{1}\\ \mathrm{S\_Iid1}^{2} & \mathrm{S\_Iid2}^{2} & \mathrm{S\_Iid3}^{2} & \cdots & \mathrm{S\_IidK}^{2}\\ \vdots & \vdots & \vdots & \ddots & \vdots\\ \mathrm{S\_Iid1}^{L} & \mathrm{S\_Iid2}^{L} & \mathrm{S\_Iid3}^{L} & \cdots & \mathrm{S\_IidK}^{L} \end{bmatrix}.

不同层搜索得到的序列可能不同,例如 S_Iid1iS_Iid1i+1\mathrm{S\_Iid1}^{i}\neq\mathrm{S\_Iid1}^{i+1}。随后即可查询其对应的缓存结果,以支持图 4(b)中的 MARM 模块。

2.3 使用 MARM 支持其他模型

工业推荐系统通常引入多个模型,逐层筛选候选物品:

  1. 检索模型: 不需要候选物品,只利用用户最新的交互物品,从完整的十亿级物品池中召回数千个物品。
  2. 级联模型: 一种小型排序模型,利用非常高效的模块,从检索模型生成的数千个候选物品中选出排名靠前的数百个。
  3. 排序模型: 广义上的排序模型是最复杂的模型,它从数百个候选物品中选出最优的几十个。

通常,这些模型采用不同的策略训练,彼此相互隔离。因此,一个有趣的假设是:排序模型的 MARM 缓存能否帮助其他检索模型或级联模型?不出所料,我们发现 MARM 结果可以无缝支持它们。

对于检索模型,由于没有候选物品集合,我们使用用户最近观看的 200 个物品查询 MARM 缓存结果,再用用户 ID 特征替代候选物品执行目标注意力,如图 4(c)所示。

对于级联模型,我们实现与公式(6)中排序模型相同的多层目标注意力过程。需要注意的是,级联模型不会再次积累 MARM 缓存,而是直接使用已经积累的 MARM 缓存,并且只使用最新的短期序列。这种方法既节省资源,又能增强级联模型与排序模型之间的一致性。

image.png 图 4:使用 MARM 框架,通过 SIM 的 GSU 与 ESU 处理长序列;以及支持检索模型的示例。

3 缓存缩放定律

本节探索一种有望解锁推荐系统未来的新缩放定律,即 MARM 的缓存缩放定律。

3.1 评估设置

3.1.1 数据集

为评估 MARM 的有效性,我们在真实推荐系统中进行了详细实验。具体而言,该短视频场景包含约 3000 万用户和 6200 万个短视频,每位用户每天平均浏览 133 个短视频。我们收集了该场景过去六个月的日志,采用流式训练方式,在朴素 MARM 设置下进行常规多任务学习,即不使用 SIM。

3.1.2 指标

为简洁起见,我们只展示长播预测任务的 GAUC 和训练损失来反映模型性能。GAUC 与在线效果的相关性最高,它对每位用户的 AUC 进行加权平均,权重由该用户的样本数量决定:

GAUC=uwuAUCu,wu=logsutotal_logs.(10)\operatorname{GAUC}=\sum_u w_u\operatorname{AUC}_u,\qquad w_u=\frac{\operatorname{logs}_u}{\operatorname{total\_logs}}.\tag{10}

我们报告最后一天的平均 GAUC 和训练损失。

3.2 缓存缩放定律讨论

在继续讨论之前,先介绍一个重要概念:缓存规模 C=LndC=Lnd,其中 LL 是 MARM 的注意力层深度,nn 是物品序列长度,dd 是表征维度。值得注意的是,缓存的计算复杂度和存储资源量都与 CC 成线性正比。这意味着 CC 的大小与推荐系统中 MARM 模块的在线延迟,以及训练、推理和存储涉及的全部资源均呈线性正相关,从而确保 FLOPs 的增加能够被缓存技术带来的效率收益抵消。

3.2.1 扩展维度 dd 与序列长度 nn

一般来说,表征维度 dd 对建模效果的影响与所用特征及数据集复杂度相关。为避免组合爆炸,我们在不使用 MARM 模块的情况下,对表征维度 dd 和序列长度 nn 进行网格搜索。此时 L=0L=0,等价于只有一层目标注意力的 DIN。如图 5(a)所示,当 dd 达到 128 后,继续增加维度对性能的影响很小。因此,在后续实验中,我们将维度固定为 d=128d=128,以进行更全面的分析。

3.2.2 扩展深度 LL 与序列长度 nn

如图 5(b)所示,增加 MARM 深度 LL 和序列长度 nn 与性能之间存在明显的正相关关系。此外,即使用户行为序列长度增加到 6400,提高 MARM 深度 LL 仍能明显改善模型性能。这证明,缓存中间结果可能是解锁更优推荐系统构建方向的一条有效路径。

3.2.3 扩展缓存规模 CC

基于上述研究,我们进一步分析模型性能与缓存规模 CC 的关系。具体而言,我们绘制多条曲线,表示缓存规模和 FLOPs 相同、但配置不同的变体。例如,对于 400×128400\times128,存在 4×100×1282×200×1281×400×1284\times100\times128\leftrightarrow2\times200\times128\leftrightarrow1\times400\times128 三种变体。

如图 5(c)所示,随着缓存规模增大,模型性能呈现明显的幂律提升趋势。一个有趣的现象是,当缓存规模较小时,增加序列长度 nn 的效果明显优于增加深度 LL,例如在 200×128200\times128 时便是如此。然而,当缓存规模达到一定水平后,增加序列长度与增加 MARM 深度的效果开始接近。由此可以得到一个令人振奋的观察:只要缓存规模相同且足够大,模型就会表现出相近的性能,例如 6400×1286400\times128

image.png 图 5:MARM 随规模扩展时的模型性能。

3.3 真实场景实验

3.3.1 对比方法详情

总体而言,MARM 可以视为用户序列建模模块,因此我们选择以下强方法来验证 MARM 的能力。

  1. Baseline: 使用多任务混合专家结构,并包含用户、视频和统计特征。用户序列信息由用户短期历史行为的求和池化结果表示。
  2. DIN: 最常用的用户短期历史行为建模算法,采用目标注意力机制。这里使用的用户历史长度为 50。
  3. SIM Soft: 采用 GSU 与 ESU 的两阶段建模方法。GSU 使用视频的预训练多模态嵌入计算内积,从用户历史中选择相关性最高的 Top-KK 视频。这里使用的用户历史长度为 15000。
  4. TWIN: SIM 架构下的两阶段建模方法,通过对齐 GSU 与 ESU 的计算方式来增强二者之间的一致性。这里使用的用户历史长度为 15000。
  5. TWIN V2: 使用层次聚类缩减超长用户历史,再通过 TWIN 方法建模。这里使用的用户历史长度为 100000。
  6. HSTU: 使用 HSTU 风格的多层自掩码注意力机制建模用户历史,时间复杂度为 O(n2d)O(n^2d)。由于计算资源有限,每个历史物品只使用少量特征,包括物品 ID、作者 ID、标签和用户反馈;历史长度为 2000,深度为 4。其 FLOPs 可以视为不使用缓存的 MARM 的 FLOPs。
  7. MARM: 使用 MARM 与 SIM 结合的方法。第一层复用现有 TWIN 模块,再堆叠 MARM 模块。每个 MARM 模块存储的序列长度为 6000,最大注意力深度 LL 为 4。

实际上,MARM 框架可以堆叠在任意既有用户序列建模模块之上。当然,我们也承认,基础模块会显著影响最终结果。

基于上述方法,我们按照学术界和工业界的惯例进行了两类实验:

  1. 独立实验: 只包含基线模型和对应的方法改造。
  2. 集成实验: 某项改造被验证有效后,就与后续改造依次融合。

3.3.2 独立实验性能比较

本节将不同序列建模模块分别引入基线模型,实验结果见表 1。

表 1:与最先进方法的离线比较。每个模块均单独添加到基线模型中。MARM 数据中的箭头表示与未使用缓存的 MARM 的 FLOPs 进行比较。最佳和次佳结果分别以粗体和下划线标出。 image.png

与所有基线方法相比,MARM 取得最佳性能。值得注意的是,MARM 没有像 TWIN V2 那样使用长度达到 100000 的用户历史,而只使用长度为 6000 的历史,但性能仍显著优于 TWIN V2。

我们还列出了每个模块的 FLOPs 指标。需要说明的是,在由 GSU 和 ESU 组成的 SIM 两阶段建模架构中,虽然训练期间 ESU 的序列长度只有几百,但为了公平,计算 FLOPs 时仍纳入此前 GSU 部分的计算成本。得益于 MARM 在深度方向上的线性可扩展性,我们不仅可以扩展用户序列长度,也可以扩展注意力深度。可以观察到,随着注意力深度增加,MARM 仍能带来显著提升。

3.3.3 集成实验的离线性能

本节旨在回答一个问题:对于已经融合多种长期和短期用户行为建模方法的模型,MARM 模块能否进一步改善性能?为此,我们进行了另一组集成实验,结果见表 2。具体来说,我们依次向基线模型中加入 DIN、SIM Soft、TWIN、TWIN V2 和 MARM 模块,并观察它们带来的提升。

表 2:与最先进方法的集成比较。若某个模块带来的增益达到显著置信水平,就将其依次加入模型。以 TWIN V2 为例,其结果包含 Baseline、DIN、SIM Soft、TWIN 和 TWIN V2。 image.png

根据表 2,在加入一系列长期和短期序列建模模块后,再加入一个简化的 HSTU 风格模块并未带来显著提升,可能是因为所用特征规模、注意力深度或历史长度尚未达到临界规模。另一个可能原因是,MARM 能够保存更多交互样本知识。纯物品序列只有几十种属性,而 MARM 的缓存结果始终是数千维的、经过压缩的目标物品查询信息,这使 MARM 内部可能发生样本间的信息交流。此外,直接加入 HSTU 风格模块会显著增加模型的计算负担。

在我们的场景中,即使系统已经包含多个长期和短期序列建模模型,加入 MARM 仍能显著提高准确率。MARM 是一个非常实用的模块,能够以高度可控的成本加入大多数推荐模型。

3.3.4 MARM 中 GSU 搜索 Top-KK 结果的重叠率

首先,我们希望理解 MARM 注意力深度为何能够持续带来提升。在 MARM 与 SIM 结合的框架中,MARM 块的每一层都像 TWIN 一样建模 GSU 和 ESU 两个阶段。GSU 会返回与目标物品最相关的 Top-KK 用户历史。

如图 6 所示,我们分析了四层 MARM 块中各层 GSU 返回历史的重叠情况,发现每一层的独立内容比例都比较高,超过 50%50\%;任意两层之间的直接重叠率则低于 20%20\%。这表明,MARM 的每一层都关注不同的历史内容,共同形成兴趣的高层次表达。此外,第 4 节还会详细讨论 MARM 缓存策略的创新性和有效性。

image.png 图 6:不同层之间 GSU Top-KK 结果的重叠率。

3.3.5 MARM 成本讨论

HSTU 可以视为未缓存 MARM 的等价形式。它采用先聚合样本、再进行学习的范式来缓解计算压力。对于拥有终身历史行为的用户,其整体计算复杂度约为 O(Ln2d(N/r))O(Ln^2d(N/r)),其中 NN 表示用户的交互总量,rnr\ll n 是用户样本的聚合级别,即每 rr 条用户样本聚合后进行一次训练。实践中,更大的聚合参数 rr 通常可以降低建模复杂度,却可能延迟用户反馈进入模型训练,容易导致在线指标下降,从而形成成本与效果之间的权衡。

相比之下,HSTU 使用 FLOPs 较高的掩码自注意力;MARM 则是资源需求更低的解决方案,无需聚合训练样本即可进行用户建模。对于一位用户的完整学习过程,MARM 的计算复杂度约为 O(Lnd)O(Lnd)。作为代价,MARM 确实需要额外的存储,即前文定义的缓存规模 CC。在我们的场景中,当注意力深度 L=4L=4、序列长度 n=6000n=6000 时,MARM 使用 60 TB 存储。存储成本与增加的计算开销之和,大约只有直接采用多层自注意力方案的八分之一。

已部署的 MARM 版本使用 L=4L=4n=6000n=6000d=128d=128,服务 3000 万日活跃用户。MARM 只使用 100 块 A10 GPU,成本约为每年 350 万元,外加 60 TB 存储,成本约为每年 120 万元;这些资源覆盖训练和推理的全部环节。最先进的 HSTU 至少需要 MARM 十倍的计算资源,即 1000 块 A10 GPU,成本约为每年 3500 万元。存储成本大约只有额外计算成本的三分之一,这正是 MARM 以存储换计算如此具有成本效益的原因。

图 7 还展示了 MARM 给推理服务带来的额外时间成本。相较于 MARM 带来的性能提升,增加 15 毫秒是值得接受的权衡:它与 TWIN 增加的 15 毫秒相当,显著优于 HSTU 增加的 40 毫秒。

image.png 图 7:已部署 MARM 版本的在线排序时间开销约为 15 毫秒,其中embedding拉取约占 5 毫秒,推理前向计算约占 10 毫秒。

3.3.6 在线结果

为量化 MARM 模型对真实推荐系统的贡献,我们在检索、级联和排序阶段分别实现了 MARM,并通过在线 A/B 测试系统验证其有效性。模型根据核心播放时长指标和交互指标进行评估,包括用户人均观看时长和点赞数等。表 3 展示了 MARM 在不同推荐阶段的在线结果。

表 3:短视频服务的在线 A/B 测试结果。 image.png

App 留存指标为正,说明 MARM 上线后带来了更好的用户体验。尤其是在多个阶段中,MARM 使用户人均 App 使用时长显著提升 2.079%2.079\%,凸显了它在增加用户观看时长方面的重要作用。虽然评论、转发和关注指标略有负向变化,但点赞交互取得了显著增益。由于推荐系统的首要目标是最大化人均观看时长,在线结果仍处于合理的指标权衡范围内。

MARM 已经在真实推荐系统中取得显著业务收益,并持续为数千万用户服务超过一年。经过不断迭代,MARM 已为业务带来超过 5%5\% 的用户 App 使用时长提升。

4 缓存策略讨论

4.1 创新说明

自然语言处理中的 KV 缓存技术通常应用于自回归大语言模型的推理阶段,用来降低生成下一词元的时间复杂度。MARM 缓存技术则同时用于训练阶段和推理阶段。总体而言,我们进行了以下富有启发性且至关重要的创新。

  1. 缓存生成。 自然语言处理中的 KV 缓存会在模型推理和预测下一词元时生成缓存结果。MARM 则在训练过程中生成缓存结果,目的是保存所有用户与物品之间的行为模式信息。在 MARM 推理期间,模型会共享训练生成的缓存结果,并利用这些已保存结果作出更好的预测。
  2. 缓存生命周期。 自然语言处理中的 KV 缓存只服务于当前句子的生成,因此只需临时存储在 GPU 显存中。句子生成完毕后,对应缓存数据便不再需要,可以被删除。与具有明确终点的语言句子不同,流式推荐系统中的用户与物品行为会依次不断产生,用户没有代表行为结束的“终点”。因此,MARM 缓存应具有长期生命周期,用来存储用户兴趣,确保模型可以随时访问缓存结果,提供高质量推荐。
  3. 缓存范围。 自然语言处理中的 KV 缓存会保存所有词元变换后的键和值,等待后续到达的查询对其进行聚合。MARM 不缓存变换后的键和值,而是缓存最终计算得到的查询输出,从而降低整体存储需求并有效节省资源。
  4. 复杂度与资源。 在自然语言处理和推荐领域,HSTU 等类似生成式模型的训练与推理计算复杂度都是 O(n2)O(n^2)。尽管这些模型在推理时会采用 KV 缓存等技术,但仍需要至少完整计算一次整个序列的自注意力。MARM 缓存序列中每一层的计算结果,因此训练与推理的计算复杂度均为 O(n)O(n),但也会额外占用外部存储资源。

4.2 有效性分析

与 HSTU 等使用完整 Transformer 结构的计算方法相比,基于缓存的 MARM 在计算意义上基本等价,主要存在以下两点差异。

  1. MARM 每一层的键值对都从缓存中读取,因此相邻两个注意力层之间的键值对不再作为节点连接在计算图中,只有查询彼此连接。这意味着在反向传播过程中,梯度不会在这些键值对之间传递,只会从后一层的查询部分传递到前一层的键值部分。
  2. MARM 中的缓存只更新一次,即某个特定样本参与训练、完成计算并写入缓存时。写入缓存的部分会被冻结,而计算图中保留的模型参数,例如每层的 Q、K、V 映射矩阵和 FFN 参数,则会继续更新。

下面解释这两项差异为何不会明显损害 MARM 架构的性能。在流式推荐系统中,已经向特定用户曝光的物品不会再次曝光。因此,推荐场景通常使用单轮流式数据进行训练;反复训练同一个样本通常无法改善性能,甚至可能损害推荐效果。由此,推荐模型在收敛后仍会保持较强的泛化能力,能够处理新出现的物品并准确建模其特征。

对于第一项差异,如果模型已经具有很强的泛化能力,能够准确建模用户与物品对,那么计算结果就可以直接存储和使用,无需在后续训练中进行大幅调整,只需调整仍保留在计算图中的部分参数即可。

对于第二项差异,我们对推荐模型底层参数的分析表明,这些参数变化相对缓慢,而越接近 Logit 的参数变化越快。因此,冻结底层结果不会迅速导致漂移。此外,在流式推荐场景中,每个物品通常都有自己的生命周期。在其生命周期内,使用冻结参数计算出的结果,相比使用最新模型参数未必会明显变差。这就是基于缓存的 MARM 方法在流式推荐场景中仍能表现出优异性能和可扩展性的原因。

5 结论

本文提出 MARM,这是一种利用内存缓存加速推理的多层推荐模型。在推荐场景中,计算复杂度是一项重要的性能约束。我们通过缓存保存复杂模型的部分计算结果,将单次建模的复杂度从 O(n2d)O(n^2d) 降低为 O(nd)O(nd)

基于 MARM 框架,我们能够以线性资源消耗将序列建模从单层目标注意力扩展至多层,显著突破推荐模型的计算瓶颈,并支持对用户终身历史进行建模。我们探索了 MARM 框架中的缩放定律,证实缓存规模与推荐性能之间存在正向的规律性关系。此外,MARM 方法可以无缝集成到精排、粗排和检索等现有推荐模型中。

6 生成式 AI 使用声明

本文只使用 AI 工具修正语法错误。研究动机、方法和实验结果完全来自真实业务场景中的一手实验与分析。所有数据和观察结果均经过在线 A/B 测试与离线分析的严格验证。