Tree Training: Efficient LLM Training on Tree-Structured Trajectories
Abstract
Whenever LLM training involves multiple outputs from the same input---agentic multi-turn trajectories, self-consistency reasoning, beam search distillation, or multi-rollout SFT/RL---the tokens form a tree-structured trajectory with shared prefixes. Existing pipelines linearize such data and treat each branch independently, causing substantial redundant computation that scales with the number of samples. We derive that averaging the loss over all branches is algebraically identical to a per-token weighted loss, reducing the problem to computing each token's log-probability exactly once. We propose DFS serialization of the tree, which visits every token exactly once, and adapt full-attention and SSM layers to ensure the resulting log-probabilities match independent per-branch computation exactly. For memory-constrained settings where the full tree exceeds GPU capacity, we propose Redundancy-Free Tree Partitioning, which achieves zero redundant computation with peak memory bounded by a single root-to-leaf path. Together, these contributions form Tree Training, achieving up to 6.2× end-to-end training speedup on dense and MoE models for both supervised fine-tuning and reinforcement learning.