Tree Training: Accelerating Agentic LLMs Training via Shared Prefix Reuse

发表时间: 2025-11 · arXiv:2511.00413

原文: https://arxiv.org/abs/2511.00413

Shaojie Wang * † 1 Jinghui Wang * † 1 Yinghan Cui * 1 Xuxing Chen * 1 Chao Wang * 1 Liang Huang 1 Xiaojiang Zhang 1 Junyi Peng 1 Li Wan 1 Haotian Zhang 1 Bin Chen 1

速读

一句话结论 本文针对代理式大语言模型训练中多轮交互产生的树状轨迹,提出了一种基于梯度恢复和树打包的 Tree Training 框架,在不损失模型精度的前提下消除了共享前缀的冗余计算,实现了高达 6.2 倍的端到端训练加速。

要解决什么问题 代理式大语言模型(Agentic LLM)在训练时,因并发工具调用、思维模式(Think-mode)等机制,其轨迹天然呈现带有大量共享前缀的树状结构。现有做法是将树状轨迹线性化,把每条从根到叶的路径拆解为独立序列训练。其卡点在于前向与后向传播的因果不对称性导致的巨大计算冗余。前向传播中因果掩码使注意力矩阵呈下三角,相同前缀在不同分支输出相同,让推理时的前缀缓存(Prefix Caching)成为可能。但在后向传播中,值梯度 $dV = P^T \times dO$,转置使 $P^T$ 变为反因果的上三角矩阵。这意味着前缀 Token 的梯度必须聚合其后所有后缀 Token 的梯度贡献。即使前缀相同,只要后缀不同,回传给前缀的梯度就不同。若在训练中强行缓存,必须在显存中存储所有分支的后缀信息,直接导致显存爆炸。因此现有框架只能让每个分支独立重算共享前缀,在面对大量重叠前缀的数据时造成严重算力浪费。

怎么做的 核心思路是摒弃线性化拆解,采用树状序列化输入,让共享前缀仅保留一份并被所有子节点复用。为绕开后缀状态导致显存爆炸的卡点,作者提出梯度恢复(Gradient Restoration)方案。其洞察在于梯度传递具线性可加性,只要保证前缀在树训练中获得的输入梯度,等于基线方法中所有对应前缀分支的梯度之和,就能实现参数更新的严格等价。关键不变量为: $$d Y_{P}^{ours} = \sum_{i=1}^{n} d Y_{p_{i}}^{base}$$ 基于此理论,系统由四个关键部件构成。一是共享前缀注意力掩码(Shared Prefix Attention Mask),它修改标准因果掩码,限制 Token 的注意力可见范围,确保不同分支能安全读取共享前缀且彼此隔离,防止信息泄露。二是位置嵌入(Position Embedding)恢复模块,树状数据打包会打乱物理位置,该模块强制为每个 Token 恢复其在原始树结构中的真实 Position ID,保证旋转位置编码(RoPE)等依赖位置的操作计算无误。三是梯度缩放器(Gradient Scaler),这是核心部件。它提前计算每个前缀节点被下游轨迹复用的次数。在反向传播前,将该节点的初始梯度乘以复用次数。这在数学上等价于将前缀独立计算多次并累加梯度,彻底免去存储多份后缀状态的显存开销。四是树打包(Tree Packing)策略,针对单棵树可能撑爆显存的问题,采用基于启发式深度优先搜索(DFS)的分区算法,优先分配最深叶节点并将深度相似节点分组,在满足显存上限时将大树切分为多个子树序列,最大化前缀重用率。

效果如何 实验在 64 张 NVIDIA Hopper GPU 集群上基于 Megatron-Core 框架搭建。评测模型为 32B 参数密集模型(Qwen3-32B)和 30B 参数混合专家模型(Qwen3-30B MoE)。数据包含来自 Terminus 和 Claude code 的真实多轮强化学习轨迹,以及重叠率 20% 到 92% 的合成数据。对比基线是标准的 Sequence Packing 策略,代表将树状数据完全展平、线性化拼接的传统路线。量化结果显示,在真实数据上,Tree Training 在 32B 密集模型实现 6.3 倍端到端加速,在 30B MoE 模型实现 6.2 倍加速,达理论上限 95% 以上。在合成数据高重叠率理想设置下,加速比最高达 8.7 倍。正确性方面,训练 Loss 曲线与基线完全重合,相对误差小于 1%。在 Terminal Bench 2.0 任务中,全树数据训练使模型得分从 20.9 提升至 28.8。该方法代价极低,掩码、位置 ID 和缩放器带来的额外显存开销不到 1MB,相比基线 64GB 的激活显存需求可忽略不计。局限在于,当单棵树过大必须依赖树打包切分时,加速比会略低于全树放入显存的理想状态。

A1 主要贡献

本文针对Agentic LLM(代理式大语言模型)训练中普遍存在的多轮交互和分支路径问题,提出了一种高效的训练框架——Tree Training。主要贡献如下:

  1. 核心问题识别:作者指出Agentic LLM的训练轨迹(Trajectory)由于并发工具调用、Think-mode(思维模式)、子代理(Sub-agents)等设计,天然形成带有共享前缀的树状结构(Tree-structured),而非简单的线性序列。现有的训练流程通常将这些轨迹线性化并独立处理每个分支,导致前向和后向传播中存在大量冗余计算。
  2. Gradient Restoration(梯度恢复):这是本文的核心创新点。不同于仅适用于推理的KV Cache,作者提出了一种在后向传播中消除冗余前缀计算的方法。该方法允许每个共享前缀仅计算一次,通过对梯度进行特定的补偿(Scaling),使得最终聚合的梯度在数学上严格等价于对所有分支独立训练的结果,且开销极低。
  3. Tree Packing(树打包):为了解决单个轨迹树可能超过GPU显存限制的问题,作者重新设计了训练引擎以支持树状数据输入,并提出了一种基于启发式DFS的内存高效分区策略(Tree Packing)。该策略将大树分割为多个子树,在满足显存约束的同时最大化前缀重用率。
  4. 显著的性能提升:在密集模型(Dense)和混合专家模型(MoE)上的实验表明,该方法在监督微调(SFT)和强化学习(RL)的模型更新阶段,均能实现高达 6.2倍 的端到端训练加速,且不损失模型精度。
共享前缀示意图。展示了一个思维模型执行多轮任务的示例:三个序列分别对应第1、2、3轮的模型输入和响应。每一轮新的思维过程会丢弃上一轮的思维过程,导致前一轮输入输出的拼接不等同于当前轮的输入上下文。因此,任务完成后,所有生成的token形成了一个树状结构。将这些共享大量重叠前缀的序列独立训练会导致显著的冗余计算。
共享前缀示意图。展示了一个思维模型执行多轮任务的示例:三个序列分别对应第1、2、3轮的模型输入和响应。每一轮新的思维过程会丢弃上一轮的思维过程,导致前一轮输入输出的拼接不等同于当前轮的输入上下文。因此,任务完成后,所有生成的token形成了一个树状结构。将这些共享大量重叠前缀的序列独立训练会导致显著的冗余计算。

A3 背景知识/关键Observation/设计原则

前向与后向传播的不对称性
在自回归(Autoregressive)LLM中,前向传播(Forward Pass)利用因果掩码(Causal Mask),即输出 $O_i$ 仅依赖于当前及之前的Token。这使得推理时的 Prefix Caching(前缀缓存)成为可能:相同的字首(Prefix)产生相同的Key/Value状态,可以被复用。

然而,训练时的后向传播(Backward Pass)呈现“转置”的因果关系:

$$\begin{aligned} \begin{aligned} S &= Q \times K^{T} \\ P &= softmax(S) \\ O &= P \times V \end{aligned} \end{aligned}$$


由于因果掩码,注意力矩阵 $P$ 是下三角矩阵。相同的前缀在不同序列中具有相同的 $P$ 和 $V$,因此输出 $O$ 相同。

$$dV = P^T \times dO$$
矩阵转置导致 $P^T$ 变为上三角矩阵(反因果)。这意味着前缀Token的梯度 $dV_i$ 需要聚合来自该Token之后所有后缀Token的梯度贡献。

关键Observation:即使两个序列拥有完全相同的前缀,由于后缀不同,反向传播回前缀的梯度 $dV_{prefix}$ 也会不同(如图2所示)。传统的缓存方法若要支持训练,必须存储所有分支的后缀信息,这将导致显存消耗爆炸,在实践中不可行。因此,必须寻找一种数学上等价但无需存储所有后缀状态的方法来聚合梯度。

树状数据集的数据预处理、前向传播和后向传播示意图。粉色块代表前缀,黄色块代表两个不同的后缀。(Q1, K1, V1, O1, dO1, dV1) 对应共享前缀的查询、键、值、输出、输出梯度和值梯度矩阵。在前向传播中,O1 和 O1' 是相同的,使得缓存可行。在后向传播中,中间变量 dV1 和 dV1' 不相同,使得直接缓存不可行。
树状数据集的数据预处理、前向传播和后向传播示意图。粉色块代表前缀,黄色块代表两个不同的后缀。(Q1, K1, V1, O1, dO1, dV1) 对应共享前缀的查询、键、值、输出、输出梯度和值梯度矩阵。在前向传播中,O1 和 O1' 是相同的,使得缓存可行。在后向传播中,中间变量 dV1 和 dV1' 不相同,使得直接缓存不可行。

A2 方法细节

3.2 梯度恢复 (Gradient Restoration)

基线方法 vs. 树训练方法

$$\begin{aligned} \begin{aligned} X_{base}(i) = \text{Concat}[ & \text{token}(i), X_{base}(\text{child}(0, i)), \\ & \text{token}(i), X_{base}(\text{child}(1, i)), \dots, \\ & \text{token}(i), X_{base}(\text{child}(m, i))] \end{aligned} \end{aligned}$$

$$X_{ours} = [P; S_1; S_2; ...; S_n]$$

梯度更新的等价性推导
为了证明树训练的有效性,必须保证两种方法计算出的梯度是等价的。考虑一个根节点带有 $n$ 个子节点的简化2层子树情况。
对于线性变换 $Y = X \times weight$,权重梯度为 $dweight = X^T \times dY$。基线方法将前缀 $P$ 与每个后缀 $S_i$ 拼接,而本文方法仅保留一个 $P$。
为了使梯度贡献等价,必须满足:

$$P^{T} \times d Y_{P}^{o u r s}=P^{T} \times(\sum_{i=1}^{n} d Y_{p_{i}}^{b a s e})$$


这意味着,前缀 $P$ 在我们方法中对 $dweight$ 的贡献,必须等于它在基线方法中所有对应前缀贡献的总和。

后向 V-梯度计算与基线的对比说明。粉色块代表前缀部分的 (Q, K, dO, dV),黄色块代表后缀部分。橙色块表示其对应的 (S 或 P) 值在注意力计算中是激活的,白色块表示被掩码屏蔽,绿色块代表在我们方法中可以省略的 (S 或 P) 计算。
后向 V-梯度计算与基线的对比说明。粉色块代表前缀部分的 (Q, K, dO, dV),黄色块代表后缀部分。橙色块表示其对应的 (S 或 P) 值在注意力计算中是激活的,白色块表示被掩码屏蔽,绿色块代表在我们方法中可以省略的 (S 或 P) 计算。

算法分析与证明
为了保证参数更新的等价性,必须满足两个条件:

  1. 梯度聚合等价
$$P^{T} \times d Y_{P}^{ours} = P^{T} \times (\sum_{i=1}^{n} d Y_{p_{i}}^{base})$$

$$\sum_{i=1}^{n} S_{i}^{T} \times d Y_{S_{i}}^{ours} = \sum_{i=1}^{n} S_{i}^{T} \times d Y_{S_{i}}^{base}, \forall i \in [1, n]$$
2. 输入梯度等价
$$\begin{aligned} \begin{aligned} d Y_{P}^{o u r s} &=d Y_{p_{1}}^{b a s e}+d Y_{p_{2}}^{b a s e}+...+d Y_{p_{n}}^{b a s e} \\ d Y_{S_{i}}^{o u r s} &=d Y_{S_{i}}^{b a s e}, \forall i \in[1, n] \end{aligned} \end{aligned}$$

针对Transformer中的不同操作,作者分别进行了论证:

$$Y = X \times weight$$

$$dX = dY \times weight^T$$
$$\begin{aligned} \begin{aligned} dX_{P}^{ours} &= dX_{p_{1}}^{base} + dX_{p_{2}}^{base} + ... + dX_{p_{n}}^{base} \\ &dX_{S_{i}}^{ours} = dX_{S_{i}}^{base}, \forall i \in [1, n] \end{aligned} \end{aligned}$$

实现细节
基于上述理论,实现包括三个关键组件(如图6所示):

  1. 共享前缀注意力掩码 (Shared Prefix Attention Mask):在前向传播中,引入一种修改后的因果掩码。该掩码限制每个Token的注意力范围,确保不同轨迹的Token可以安全地共享前缀表示,而不会发生信息泄露(即不同分支间不可见)。作者基于 Flash Attention V3 【Shah et al., Flashattention-3: Fast and accurate attention with asynchrony and low-precision, 2024, arXiv】 实现了支持节点级共享前缀掩码的高性能GPU内核。
  2. 位置嵌入 (Position Embedding):树打包后的数据物理位置发生了变化。为了保持一致性,必须恢复原始(打包前)的Position IDs。例如,某分支的Token在树结构中可能紧跟在前缀之后,其Position ID应接续前缀的ID,而非其在打包内存中的索引。
  3. 梯度缩放器 (Gradient Scaler):这是最核心的组件。作者计算每个节点在树中被复用的次数(称为 tree-scale)。在反向传播开始前,将每个节点的梯度乘以对应的 tree-scale 因子。
    • 工作流:如图7所示,如果一个前缀节点被5条轨迹复用,其梯度会被乘以5。这在数学上等价于将该前缀独立计算5次并累加梯度。
    • 并行性:Tree Training与现有的并行策略(TP/EP/DP/PP)正交,可无缝结合。对于上下文并行(Context Parallelism),只需根据查询分片生成对应的注意力掩码即可。
    • MoE负载均衡:对于有辅助损失(Auxiliary Loss)的MoE模型(如Qwen3 MoE),通过在计算Router辅助损失时将前缀Token乘以其共享计数(Gradient Scaler),即可实现与基线的数学等价。
扁平化树状轨迹的实现细节。每个扁平化的树状轨迹需要:(1) 用于前缀复用的梯度缩放张量,(2) 恢复原始Token位置的位置嵌入张量,以及 (3) 共享前缀注意力掩码,用于在前向和后向传播期间正确复用重叠前缀的计算。
扁平化树状轨迹的实现细节。每个扁平化的树状轨迹需要:(1) 用于前缀复用的梯度缩放张量,(2) 恢复原始Token位置的位置嵌入张量,以及 (3) 共享前缀注意力掩码,用于在前向和后向传播期间正确复用重叠前缀的计算。
Tree Training中带有梯度缩放的后向计算。在计算扁平化树状轨迹的初始梯度 (dY) 后,应用梯度缩放器将共享前缀的梯度按其复用次数进行缩放。例如,共享前缀 r -> u 被5条轨迹使用(scale=5),v1 被3条轨迹使用(scale=3),从而确保梯度的正确累加。
Tree Training中带有梯度缩放的后向计算。在计算扁平化树状轨迹的初始梯度 (dY) 后,应用梯度缩放器将共享前缀的梯度按其复用次数进行缩放。例如,共享前缀 r -> u 被5条轨迹使用(scale=5),v1 被3条轨迹使用(scale=3),从而确保梯度的正确累加。

3.3 树打包 (Tree Packing)

问题与策略
在实际训练中,完整的轨迹树可能过大无法放入GPU显存。因此,需要一种打包算法将大计算树分割为一系列满足显存限制 $C$ 的子树,同时最大化前缀共享。
理论上的最优划分需要动态规划结合装箱问题(Bin Packing),属于NP-hard问题,对于大规模树计算成本过高。

启发式DFS算法
作者采用了一种贪心启发式DFS(深度优先搜索)算法,该算法随树的大小线性扩展,能有效逼近最优解。其原则包括:

  1. 优先分配最深叶节点:因为它们对总轨迹长度贡献最大。
  2. 分组:将同一子树中深度相似的叶节点组合在一起,提高打包的同质性。
  3. DFS遍历:按深度优先顺序遍历树,一旦累积长度超过容量 $C$,则启动新的遍历(生成一个新的Packed Sequence)。

效果示例:如图4所示,一个包含4条轨迹、总计83k tokens的树,若限制显存为60k tokens。基线线性化方法会产生164k tokens。而Tree Packing将其分为两个序列,总计仅102k tokens,显著减少了冗余。

显存受限下的Tree Packing示意图。理想情况下,Tree Packing将所有轨迹合并为单个序列以最大化前缀共享。然而,当总树大小(如83k tokens)超过GPU显存限制(如60k tokens)时,我们的方法将树分割为最小的子序列。与基线扁平化方法(164k tokens)相比,我们的方法(102k tokens)显著减少了冗余。
显存受限下的Tree Packing示意图。理想情况下,Tree Packing将所有轨迹合并为单个序列以最大化前缀共享。然而,当总树大小(如83k tokens)超过GPU显存限制(如60k tokens)时,我们的方法将树分割为最小的子序列。与基线扁平化方法(164k tokens)相比,我们的方法(102k tokens)显著减少了冗余。

A4 实验环境

A4 实验结果

1. 性能指标定义

2. 真实场景下的加速与正确性

真实Rollout数据的端到端训练加速和Loss对比。左图为MoE模型Qwen3-30B,右图为Dense模型Qwen3-32B。上方展示了Tree Training带来的训练加速倍数,下方展示了Loss的平均相对误差。
真实Rollout数据的端到端训练加速和Loss对比。左图为MoE模型Qwen3-30B,右图为Dense模型Qwen3-32B。上方展示了Tree Training带来的训练加速倍数,下方展示了Loss的平均相对误差。

3. 不同POR下的加速表现

Tree Training在不同POR数据集上的端到端训练加速。每个子图报告了Tree Training相对于基线的总训练时间减少比例。(a) 全树适合GPU显存的合成数据集,(b) 需要Tree Packing的显存受限合成数据集。
Tree Training在不同POR数据集上的端到端训练加速。每个子图报告了Tree Training相对于基线的总训练时间减少比例。(a) 全树适合GPU显存的合成数据集,(b) 需要Tree Packing的显存受限合成数据集。

4. 显存开销

5. 下游任务性能提升

A5 结论

本文提出的 Tree Training 框架通过 Gradient RestorationTree Packing 技术,成功解决了Agentic LLM训练中因树状轨迹线性化导致的计算冗余问题。该方法在数学上严格保证了梯度更新的正确性,即与独立训练所有分支完全等价。实验证明,该框架在真实世界的Agentic RL和SFT任务中,能够在极低的额外开销下实现显著的训练加速(最高达6.2倍),且适用于Dense和MoE等多种模型架构。这不仅提升了训练效率,也为未来利用更复杂的树状思维链和多路径交互数据进行模型训练铺平了道路。

A6 附录

附录 A: Tree Packing 动态规划 (DP) 解决方案

A.1 单路径打包 (Single-Path Packing)
首先考虑简化情况,即每个训练步仅建立一条共享路径 $[r \to u]$。

$$L(u)+R(u) \leq C.$$

$$\begin{aligned} DP(u) = \begin{cases} 0, & \text{if } u \text{ is a leaf,} \\ \max \left\{ \begin{aligned} &\mathbf{1}_{\mathrm{f}(u)} \cdot (n_u - 1)L(u), \\ &\textstyle\sum_{v \in \mathrm{child}(u)} DP(v) \end{aligned} \right\}, & \text{otherwise.} \end{cases} \end{aligned}$$
即在当前节点 $u$ 作为共享节点,或者递归地由子节点处理中取最大值。

A.2 多路径打包 (Multi-Path Packing)
单路径策略在容量 $C$ 较大时可能无法实现最大复用(如图10所示,分两步分别复用 $r \to u \to v_1$ 和 $r \to u \to v_5$ 不如一次性复用 $r \to u$ 并分叉更优)。因此扩展到多路径设置:

单路径与多路径Tree Packing的对比。步骤1打包共享前缀 r -> u -> v1,步骤2单独打包 r -> u -> v5。而最优策略是将 r -> u -> {v1, v5} 视为分层的共享前缀,从而实现更高的计算复用,如章节 A.2 所述。
单路径与多路径Tree Packing的对比。步骤1打包共享前缀 r -> u -> v1,步骤2单独打包 r -> u -> v5。而最优策略是将 r -> u -> {v1, v5} 视为分层的共享前缀,从而实现更高的计算复用,如章节 A.2 所述。