🤖 AI Summary
This study addresses the inefficiency of redundant KV computation for shared prompts during reinforcement learning post-training of long-context large language models. To this end, we propose AugTree, an execution framework that introduces a decoupled context parallelism mechanism, separating the replay phase into prompt prefilling and response branch scheduling. An online planner dynamically selects communication semantics to achieve load balancing, while differentiable KV state propagation is supported to ensure training semantic precision. By integrating KV cache reuse with partial softmax reduction, AugTree achieves up to a 7.08× speedup over baselines on 64 accelerators, significantly reducing per-step training time.
📝 Abstract
Training long-context LLM policies with RL requires re-evaluating groups of sampled responses under the updated policy, an update-stage attention workload that differs sharply from pre-training: each group shares one long prompt that fans out into multiple response branches. Standard context parallelism (CP) flattens each prompt--response pair into a linear sequence, so the same prompt key--value (KV) states are recomputed---or repeatedly rotated through the network---once per response branch. We present \textbf{AugTree}, a CP execution scheme built around this replay stage. AugTree separates the replay into two phases: a prompt-prefill phase that computes the shared prompt KV state once, and a response-replay phase that schedules the independent response branches over a bounded set of replay lanes. The replay phase instantiates two communication semantics, chosen by a lightweight online planner that enumerates CP degrees, schedules, and placements before GPU dispatch: rotating KV shards within response-local lanes when responses dominate, and moving response queries to stationary prompt-KV owners with a partial-softmax reduction when prompts dominate. The shared prompt state remains fully differentiable---response losses backpropagate into it and the accumulated prompt gradients propagate through the original prefill graph---so AugTree preserves exact training semantics rather than performing detached, inference-style KV caching. On four real post-training workloads and up to 64 accelerators, AugTree improves average training-stage step time by 1.18$\times$ over dynamic CP (up to 2.23$\times$), 2.63$\times$ over a Megatron ring CP baseline with prompt reuse, and 7.08$\times$ over the baseline without reuse.