Splitting Prompt Prefill from Response Replay for Context-Parallel Long-Context LLM Post-Training

📅 2026-09-27
📈 Citations: 0
✨ Influential: 0
📄 PDF
🤖 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.
Problem

Research questions and friction points this paper is trying to address.

long-context LLM
context parallelism
post-training
reinforcement learning
KV recomputation
Innovation

Methods, ideas, or system contributions that make the work stand out.

Context Parallelism
Long-Context LLM
Reinforcement Learning Post-Training
Prompt-Response Decoupling
Differentiable KV Replay
🔎 Similar Papers
No similar papers found.
Y
Yubing Bao
Fudan University
Z
Zhihui Lu
Fudan University
Q
Qiang Duan
The Pennsylvania State University
Yuedong Xu
Yuedong Xu
Professor, Fudan University
Machine Learning SystemsNetwork EconomicsNetwork ModelingMultimedia Networking
S
Sen Liu
Fudan University
P
Pan Zhou
Singapore Management University