← 返回论文列表
📄 生成式推荐 · 训练-推理对齐

APAO:把 Beam Search 的"剪枝压力"提前灌进训练里

Adaptive Prefix-Aware Optimization for Generative Recommendation(KDD'26)

作者
Yuanqing Yu, Yifan Wang, Weizhi Ma, Zhiqiang Guo, Min Zhang
机构
清华大学 DCST / AIR · 鹊承实验室
来源
arXiv:2603.02730v3 · KDD 2026
核心命题
生成式推荐 token-level CE 损失与 beam search 推理之间存在 prefix-level 不一致;用 prefix-aware + adaptive worst-prefix 损失对齐

📎 论文链接arxiv.org/abs/2603.02730

💻 代码github.com/yuyq18/APAO

🚀 工业落地:微信公众号平台,pCTR +0.9%,图文 CTR +0.706%

TL;DR

一句话总结:生成式推荐用 CE 训练 + beam search 推理,但 CE 假设"前缀永远正确",而 beam search 会因为前缀分数低就早早把候选项剪掉。APAO 在训练阶段把"每一个前缀都必须存活"的约束显式建模为前缀级 loss,再用自适应权重把训练精力倾斜到当下最薄弱的那个前缀,从而把训练目标真正对齐到推理过程,在 4 个公开数据集 + 微信公众号线上 A/B 都拿到稳定收益。

核心痛点

CE 优化的是 token 平均似然,可以让早期 token "差一点"被后期 token "补回来";但 beam search 是步步剪枝,早期前缀掉出 Top-K 就永远丢了

核心思路

把"每个前缀长度 m 上的排名约束"显式写成一个 loss,组合多个 m 的 loss,并按"哪个前缀此刻最菜"动态调权(指数加权 mirror descent)。

📚1. 背景与动机

1.1 什么是生成式推荐(GR)

传统判别式推荐对每个 item 学一个 ID embedding,给定用户历史,模型直接打分→排序。生成式推荐(Generative Recommendation, GR)的玩法不一样:

  • 先用 RQ-VAE / RQ-K-means 这类向量量化方法,把每个 item 切成 T 个离散的 semantic ID token(如 4 层 codebook,每个 item 表示为 4 个 token)。
  • 把用户行为序列看成一个 token 序列,训练一个 Transformer(TIGER 是 encoder-decoder,Llama 是 decoder-only),让它自回归地预测下一个 item 的 4 个 token
  • 推理时用 beam search 把 Top-K 个高分 token 序列解出来,每个序列对应一个 item。
💡 举例:一个具体的生成流程

假设用户历史是 [运动鞋, 篮球, 护腕],每个 item 已 tokenize 为 4 个 SID,比如 "运动鞋 = (3,7,12,5)"。把所有历史 item 的 token 拼接给模型,模型自回归生成下一个 item 的 4 个 token。

第 1 步:在词表 V 中给所有 candidate token 打分,beam=20 表示保留前 20 个候选前缀。第 2 步:每个前缀再扩展,又只保留前 20 个累积分数最高的"前缀+新 token"组合。直到第 4 步生成完整 item。

关键点:每一步都会"砍掉"分数不够的前缀,即使这个前缀对应的目标 item 整体很合适。

1.2 为什么不直接 Full Item Sorting?

理想方案是对整个 item 库 Y 中的每个 item 都算一个总 log-probability,再排 Top-K。但当 |Y| 上百万、上亿时,这个全空间打分代价爆炸。Beam search 用"步步取 Top-K 前缀"近似全空间排序,把复杂度从 O(|Y|) 降到 O(K·|V|·T),是生成式推荐能落地的关键。

Figure 1: Analysis of Beam Search
Figure 1(论文原图):(a) Beam search 流程;(b) 推理耗时对比;(c) 黄色 baseline 中有大量 Global Top-20 item 在中途被剪掉,APAO(绿色)把它们留下来了——这正是论文要解决的问题。

1.3 训练-推理不一致的本质矛盾

关键观察:训练阶段用 teacher forcing——每步喂的都是 ground-truth 前缀;模型从未被"惩罚"过"我的某个前缀掉出 Top-K"。但推理阶段,模型必须自己活下来,前缀掉一次就没了。

🔍2. 训练-推理不一致的形式化

2.1 全空间排序视角(理想)

对于目标 item y,定义总分:

$$ S(y \mid x) = \sum_{t=1}^{T} \log P(y^t \mid y^1,\ldots,y^{t-1}, x) $$
符号说明
  • $y$:目标 item,被 tokenize 为 $T$ 个 token $\{y^1, \ldots, y^T\}$
  • $y^1,\ldots,y^{t-1}$:前 $t-1$ 个 token 构成的前缀
  • $x$:用户历史 token 序列
  • $P(y^t \mid y^1,\ldots,y^{t-1},x)$:模型给出的条件概率

Recall@K 的成功条件(Full-Sorting):

$$ \mathbb{I}_{\text{Full}}(y) = \mathbb{I}\bigl(\mathrm{rank}_{\mathcal{Y}}(S(y|x)) \le K\bigr) $$
符号说明
  • $\mathrm{rank}_\mathcal{Y}(\cdot)$:在整个 item 库 $\mathcal{Y}$ 中的排名
  • $\mathbb{I}$:指示函数(条件成立为 1)

关键:只看最终总分。中间某步概率低没关系,只要后面的 token 补回来就行。

2.2 Beam Search 视角(真实推理)

Beam search 的"成功条件"是一连串逐步排名约束的"且"

$$ \mathbb{I}_{\text{Beam}}(y) = \prod_{t=1}^{T} \mathbb{I}\bigl(\mathrm{rank}_t(S_t(y^{1:t})) \le K\bigr) $$
符号说明
  • $S_t(y^{1:t}) = \sum_{j=1}^{t} \log P(y^j \mid y^1,\ldots,y^{j-1}, x)$:长度为 $t$ 的前缀的累积分
  • $\mathrm{rank}_t(\cdot)$:第 $t$ 步在 beam 候选中的排名
  • 连乘 $\prod_{t=1}^{T}$:任何一步排名 $> K$,整个 item 就被"判死刑"
本质冲突:CE loss 优化的是 token-level 平均似然(允许早晚 token 互相补偿);beam search 要求每一步都进 Top-K,对中间失败零容忍。Teacher forcing 训出来的模型从来没"练过"在弱前缀上的鲁棒性。
💡 举例:CE vs Beam Search 的偏好差异

假设两个候选 item 的逐 token log-prob 如下(target item 是 A):

· Item A:[-3.0, -0.5, -0.5, -0.5] → 总分 -4.5

· Item B:[-1.0, -1.5, -1.5, -1.5] → 总分 -5.5

👉 Full Sorting:A 排第一(总分更高),完美。

👉 Beam Search (K=3,假设第 1 个 token A 的 -3.0 排第 10):A 在第一步就被剪掉,永远拿不回来。最终 B 反而被推荐了!

CE loss 看不到这种问题——它觉得 A 的总和最优就好。APAO 通过 prefix-aware loss 提醒模型:"第一个 token 也不能太差!"

🛠️3. APAO 方法详解

3.1 整体框架:Prefix-aware Optimization

核心思路非常直接:除了在完整 item 上算 CE,额外把"每个前缀长度 m 上的 loss"也加进去,并各自配一个权重 $w_m$:

$$ \mathcal{L}_{\text{unified}} = \mathcal{L}_{\text{CE}} + \beta \sum_{m=1}^{T} w_m \mathcal{L}_m $$
符号说明
  • $\mathcal{L}_{\text{CE}}$:标准 token-level cross-entropy
  • $\mathcal{L}_m$:长度为 $m$ 的前缀上的 prefix-aware loss(pointwise 或 pairwise)
  • $w_m \ge 0$:第 $m$ 个前缀的权重(动态学习,见 3.4)
  • $\beta \ge 0$:控制 prefix loss 整体强度的超参,论文实验在 0.1–0.4 之间最佳
  • $T$:每个 item 的 token 数量(论文里是 4)
One-stage 训练:不同于 S-DPO 那种 SFT→DPO 的两阶段流程,APAO 把 CE 和 prefix loss同时优化,更简洁、训练更快、效果更稳。

3.2 Pointwise 版本(Lpoint)

最自然的写法:把 CE 套到前 m 个 token 上,让模型对"前 m 步"的 ground-truth token 概率最大化:

$$ \mathcal{L}_{\text{point}}(m) = -\frac{1}{m} \sum_{t=1}^{m} \log \frac{\exp(z^t_{i, y^t_i})}{\sum_{j \in V} \exp(z^t_{i,j})}, \quad m \in \{1,\ldots,T\} $$
符号说明
  • $z^t_{i,v} = f_\theta(v \mid y_i^1,\ldots,y_i^{t-1}, x)$:第 $t$ 步赋给词表中 token $v$ 的 logit
  • $y^t_i$:目标 item 的第 $t$ 个 ground-truth token
  • $V$:token 词表
  • $1/m$:长度归一化,避免长前缀权重过大

直觉:它和 CE 共享同样的 logits,没有任何额外 forward,所以 pointwise 几乎不增加训练成本。本质上是给"不同前缀长度"分配不同的训练 emphasis。

💡 举例:Pointwise 在干嘛

假设目标 item 的 4 个 SID = [3, 7, 12, 5]。

· $\mathcal{L}_{\text{CE}}$:让模型对 (3→7→12→5) 这个完整序列的平均似然最大。

· $\mathcal{L}_{\text{point}}(1)$:单独再算一次"给定 user history,第 1 个 token 选 3 的 log-prob",强化第一步。

· $\mathcal{L}_{\text{point}}(2)$:算前两步 (3,7) 的平均 log-prob。

四个 $\mathcal{L}_{\text{point}}$ 加权求和。本质上就是给"前缀越短"的位置更多关注机会,让早期 token 别拖后腿。

3.3 Pairwise 版本(Lpair)

Pointwise 只学绝对概率,不直接学"正样本前缀 vs 负样本前缀"的相对排名。Pairwise 把这件事补上:

$$ \mathcal{L}_{\text{pair}}(m) = -\log \sigma \!\left( -\log \sum_{j \in \mathcal{N}} \exp\bigl(S^m_{j,-} - S^m_{i,+}\bigr) \right), \quad m \in \{1,\ldots,T\} $$
符号说明
  • $S^m_{i,+} = \sum_{t=1}^{m} s^t_i$:正样本前缀的累积 log-prob
  • $S^m_{j,-}$:从语料库随机采的负样本 item 截到长度 $m$ 的前缀累积分
  • $\mathcal{N}$:负样本集合(论文中固定 100 个)
  • $\sigma$:sigmoid 函数

形式上像 S-DPO,但有本质区别:S-DPO 只在完整 item层面对比,APAO 在每个前缀长度都对比一次。这意味着模型被强迫:"在每一个 beam 步上,正样本前缀的分都要高过 100 个负样本前缀"。

💡 举例:Pairwise 的训练信号

正样本 item 是 [3,7,12,5],随机抽 100 个负样本,比如 [9,2,3,8]、[14,7,5,1] ……

· m=1 时比较:正样本第一个 token "3" 的 logit 要大于 100 个负样本第一个 token 的 logit。

· m=2 时比较:累积分 (logp(3)+logp(7|3)) 要大于 (logp(9)+logp(2|9)) 等所有负样本前两个 token 的累积分。

· …直到 m=4。

这正是 beam search 在每一步要做的事——APAO 把它提前"演练"了。

3.4 Adaptive Worst-prefix Optimization

问题:T 个前缀的权重 $w_m$ 怎么定?手工调参组合爆炸,均匀加权又不够智能。

论文观察到:beam search 的整体成功率被"最薄弱的那个前缀"卡住(一处掉链子就全完)。所以训练时应该聚焦当前最薄弱的前缀。最朴素是 Hard-Max:

$$ \mathcal{L}_{\text{worst}} = \max_{m \in \{1,\ldots,T\}} \mathcal{L}_m $$

但 mini-batch 上的 worst 抖动很大,训练会不稳。论文借鉴 Piratla et al. (ICLR 2022) 的做法,引入软权重 + KL 平滑

$$ w^{(\tau+1)} = \arg\max_{w \in \triangle^T} \sum_m w_m \mathcal{L}_m^{(\tau+1)} - \frac{1}{\eta} \mathrm{KL}(w \,\|\, w^{(\tau)}) $$
符号说明
  • $\triangle^T$:T 维概率单纯形($w_m \ge 0, \sum w_m = 1$)
  • $\eta$:步长超参,越小权重变化越平滑(论文范围 5e-6 到 1e-4)
  • $\tau$:训练步索引
  • $\mathrm{KL}(w \| w^{(\tau)})$:与上一步权重的 KL 散度,起平滑作用

由 KKT 条件可得闭式解(指数加权 / mirror descent):

$$ w^{(\tau+1)}_m = \frac{w^{(\tau)}_m \cdot \exp(\eta \mathcal{L}_m)}{\sum_{j=1}^{T} w^{(\tau)}_j \cdot \exp(\eta \mathcal{L}_j)} $$
符号说明
  • 哪个前缀 loss 大,下一步它的权重就指数地大
  • $\eta$ 起到温度作用:$\eta$ 越大,"聚焦最差前缀"越激进;$\eta$ 越小越接近均匀
  • 每步只是一次 softmax 更新,几乎零额外开销
💡 举例:Adaptive Weighting 动态过程

初始 $w = [0.25, 0.25, 0.25, 0.25]$。某步算出 $\mathcal{L} = [2.1, 0.8, 0.5, 0.4]$(说明第 1 个前缀最菜)。

经过指数更新后,$w$ 大概变成 $[0.55, 0.20, 0.13, 0.12]$——下一轮训练会把 55% 的精力放在第 1 个前缀上。

随着第 1 个前缀被练好,它的 loss 下降,权重会被重新分配到下一个最薄弱的位置。这就像专项补课:哪个最差就先补哪个,补完换下一个

3.5 理论分析

(1) 下界保证

论文证明(Appendix B):优化 prefix-aware loss 等价于优化 beam search 下 ranking 指标的一个 lower bound,所以是 ranking 的合理代理目标。

(2) 时间复杂度

Pointwise 版本与 CE 同阶 $O(B T (d^2 + |V|d))$,因为前缀 loss 复用了已有 logits,不增加 forward。Pairwise 版本与 S-DPO 同阶,但因省掉 reference model 推理,反而比 S-DPO 更快(见表 4)。

🧪4. 实验结果

4.1 数据集和 Backbone

  • 数据集:Office、Grocery、Beauty、Yelp(用户/物品规模从 5K 到 30K)。
  • Backbone:TIGER(encoder-decoder,~0.01B)+ Llama(decoder-only,~0.01B)。
  • Tokenizer:Llama-3.1-8B-Instruct 出 embedding + 4 层 RQ-K-means 量化。
  • 推理:beam size 20;评测 Recall@10/20 与 NDCG@10/20。

4.2 主表:APAO 全面打赢 CE / MSL / DPO / DMPO / S-DPO

Backbone方法Beauty R@10Beauty N@10Yelp R@10Yelp N@10
TIGERCE (baseline)0.06110.03180.03840.0201
CE→S-DPO0.06060.03100.03890.0202
APAO-Pairwise0.0639 (+4.6%)0.0339 (+6.6%)0.0412 (+5.9%)0.0218 (+5.8%)
LlamaCE (baseline)0.05160.02740.02670.0141
CE→S-DPO0.04800.02620.02620.0138
APAO-Pointwise0.0564 (+9.3%)0.0300 (+9.5%)0.0289 (+8.2%)0.0152 (+7.8%)

数据来自论文 Table 1。在 Llama 上 APAO 涨幅尤其大(最高 13.39%),说明 decoder-only 架构对训练-推理不一致更敏感,prefix-aware 收益更显著。

4.3 消融:每个模块都有用,早期前缀最关键

Figure 2: Ablation
Figure 2(论文原图):(a) 去掉自适应权重或任一 loss 都掉点;(b) 移除 Prefix 0(第一个 token 的监督)下降最严重,验证"早期前缀最关键"的直觉。

4.4 Prefix 级 Recall:beam 越走越后,APAO 涨得越多

Figure 7: Prefix-level Recall
Figure 7(论文原图):把每一步的 beam 候选拉出来看 Recall@20。越靠后的前缀(Prefix 3 = 完整 item)APAO 相对 baseline 的提升越大——说明它真的在防止"中途被剪掉"

4.5 训练效率:Pointwise 与 CE 同等开销,Pairwise 比 S-DPO 更快

类型方法Office (s/epoch)BeautyYelp
PointwiseCE15 / 45 epoch53 / 7466 / 123
APAO15 / 4554 / 8379 / 165
PairwiseS-DPO199 / 93693 / 1091175 / 159
APAO163 / 58637 / 119946 / 144

数据来自 Table 4。Pairwise APAO 不需要 reference model 的额外 forward,所以总训练时间显著少于 S-DPO。

4.6 工业落地

线上 A/B(微信公众号平台):APAO-Pointwise 对比 CE baseline 拿到 pCTR +0.9%,图文 CTR +0.706%,图文人均点击 +0.907%,图文人均曝光 +0.205%。这是论文最有说服力的一组数据——简单换 loss、不改推理代码就能上线见效。

💭5. 我的理解与启发

5.1 论文的"小切口大问题"叙事

这篇论文最值得学习的地方是问题定义:训练-推理不一致并不是新问题(NMT 里有 exposure bias,Reinforce 类做法很多),但把这个不一致精确到"prefix 级 ranking 约束"并形式化成两个公式(Eq. 8 vs Eq. 9),让方法变得非常清晰自然——既然 beam search 是 step-wise 排序,那训练就该 step-wise 监督。

5.2 与已有工作的核心差异

vs S-DPO / DMPO

它们只在完整 item 层面做 ranking 对齐;APAO 把对齐颗粒度细化到每个 prefix 长度。这是真正"匹配 beam search 行为"的关键。

vs MSL(mask 无效 token)

MSL 改的是 softmax 范围;APAO 改的是监督粒度。两者并不冲突,理论上可以叠加。

vs DPO 系列

DPO 需要 reference model 二次 forward;APAO Pairwise 直接复用 batch 内 logits,效率更高、实现更简单。

vs 推理侧方案(如约束解码)

推理侧改动会引入额外延迟;APAO 是训练侧解法,推理代码零改动,更容易上线。

5.3 可能的延伸与不足

  • 负采样还是均匀随机:论文为公平用了 uniform random,但显然 hard negative 采样会更有效。这是明显可以补的方向。
  • Listwise loss:论文出于效率没做 listwise;但在小语义码表(如 4×256)场景下其实开销可控,值得一试。
  • Beam size 与训练时 prefix 长度的解耦:训练时所有前缀都"平等参与",但实际推理 K 越大对早期前缀容忍度越高。理论上可以让"训练时关注的前缀长度"也与推理 K 相关。
  • 规模扩展:论文规模是 0.01B 级别,OneRec / TIGER 大模型规模下 APAO 是否依然显著?文中没给。
  • Sequence-level RL 对比:论文未与 GR 上的 REINFORCE / GRPO 类方法对比,缺一个直接竞品。

5.4 对快手推荐 / OneRec 体系的启发

实践建议:OneRec、商家商品召回这类用 SID 做 token + 自回归生成的系统,本质都跑在"CE 训练 + beam search 推理"这套范式上,理论上同样存在 prefix-level 不一致。APAO-Pointwise 几乎零额外训练开销,可以作为低成本即插即用的优化项试一波。重点关注:
  • 用 prefix-level Recall 而不只是 final Recall 去评估模型 — 直接看 baseline 是否真有"中途剪掉"问题
  • $\beta$ 取 0.1–0.4,$\eta$ 取 1e-5 量级开始 grid search
  • 如果有 reference model 资源,pairwise 在小语料上更稳;大语料场景下 pointwise 性价比最高

5.5 一句话总结

APAO 没有引入新的网络结构、没有新的 tokenizer,只是把推理时的剪枝约束如实写进训练目标,加上一个"哪里弱补哪里"的自适应权重——简单、便宜、对齐了真实推理过程,最后还在工业系统跑出了 +0.9% pCTR。这是一类"看清问题本质后用最小代价解决"的好工作。