⚡TL;DR
核心痛点
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),是生成式推荐能落地的关键。
1.3 训练-推理不一致的本质矛盾
🔍2. 训练-推理不一致的形式化
2.1 全空间排序视角(理想)
对于目标 item y,定义总分:
- $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):
- $\mathrm{rank}_\mathcal{Y}(\cdot)$:在整个 item 库 $\mathcal{Y}$ 中的排名
- $\mathbb{I}$:指示函数(条件成立为 1)
关键:只看最终总分。中间某步概率低没关系,只要后面的 token 补回来就行。
2.2 Beam Search 视角(真实推理)
Beam search 的"成功条件"是一连串逐步排名约束的"且":
- $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 就被"判死刑"
假设两个候选 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{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)
3.2 Pointwise 版本(Lpoint)
最自然的写法:把 CE 套到前 m 个 token 上,让模型对"前 m 步"的 ground-truth token 概率最大化:
- $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。
假设目标 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 把这件事补上:
- $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 个负样本前缀"。
正样本 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:
但 mini-batch 上的 worst 抖动很大,训练会不稳。论文借鉴 Piratla et al. (ICLR 2022) 的做法,引入软权重 + KL 平滑:
- $\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):
- 哪个前缀 loss 大,下一步它的权重就指数地大
- $\eta$ 起到温度作用:$\eta$ 越大,"聚焦最差前缀"越激进;$\eta$ 越小越接近均匀
- 每步只是一次 softmax 更新,几乎零额外开销
初始 $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@10 | Beauty N@10 | Yelp R@10 | Yelp N@10 |
|---|---|---|---|---|---|
| TIGER | CE (baseline) | 0.0611 | 0.0318 | 0.0384 | 0.0201 |
| CE→S-DPO | 0.0606 | 0.0310 | 0.0389 | 0.0202 | |
| APAO-Pairwise | 0.0639 (+4.6%) | 0.0339 (+6.6%) | 0.0412 (+5.9%) | 0.0218 (+5.8%) | |
| Llama | CE (baseline) | 0.0516 | 0.0274 | 0.0267 | 0.0141 |
| CE→S-DPO | 0.0480 | 0.0262 | 0.0262 | 0.0138 | |
| APAO-Pointwise | 0.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 消融:每个模块都有用,早期前缀最关键
4.4 Prefix 级 Recall:beam 越走越后,APAO 涨得越多
4.5 训练效率:Pointwise 与 CE 同等开销,Pairwise 比 S-DPO 更快
| 类型 | 方法 | Office (s/epoch) | Beauty | Yelp |
|---|---|---|---|---|
| Pointwise | CE | 15 / 45 epoch | 53 / 74 | 66 / 123 |
| APAO | 15 / 45 | 54 / 83 | 79 / 165 | |
| Pairwise | S-DPO | 199 / 93 | 693 / 109 | 1175 / 159 |
| APAO | 163 / 58 | 637 / 119 | 946 / 144 |
数据来自 Table 4。Pairwise APAO 不需要 reference model 的额外 forward,所以总训练时间显著少于 S-DPO。
4.6 工业落地
💭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 体系的启发
- 用 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。这是一类"看清问题本质后用最小代价解决"的好工作。