arXivDaily arXiv每日学术速递 周一至周五更新
arXiv周末暂无论文更新,休息一下吧,周末愉快~~
arXiv 2609.33133cs.DC

将提示词预填充与响应重放分离,用于上下文并行的长上下文LLM后训练

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

Yubing Bao, Zhihui Lu, Qiang Duan, Yuedong Xu, Sen Liu, Pan Zhou

首次发表
浏览论文内容

中文总结 AI 辅助

针对长上下文LLM后训练中提示词KV重复计算问题,提出AugTree上下文并行方案,分离预填充与重放,提升训练步时间1.18-7.08倍。

中文摘要 AI 辅助

使用强化学习训练长上下文LLM策略需要在更新后的策略下重新评估多组采样响应,这是一种与预训练截然不同的更新阶段注意力工作负载:每组共享一个长提示词,该提示词扇出到多个响应分支。标准上下文并行(CP)将每个提示词-响应对展平为线性序列,因此相同的提示词键值(KV)状态会为每个响应分支重新计算——或反复通过网络轮转。我们提出了AugTree,一种围绕此重放阶段构建的CP执行方案。AugTree将重放分为两个阶段:提示词预填充阶段,该阶段一次性计算共享的提示词KV状态;以及响应重放阶段,该阶段在有限的重放通道集合上调度独立的响应分支。重放阶段实例化两种通信语义,由轻量级在线规划器在GPU调度前枚举CP度数、调度和放置来选择:当响应占主导时,在响应本地通道内轮转KV分片;当提示词占主导时,将响应查询移动到固定的提示词KV所有者,并进行部分softmax归约。共享提示词状态保持完全可微——响应损失反向传播到其中,累积的提示词梯度通过原始预填充图传播——因此AugTree保留了精确的训练语义,而不是执行分离的、推理风格的KV缓存。在四个真实后训练工作负载和最多64个加速器上,AugTree相比动态CP将平均训练阶段步时间提高了1.18倍(最高2.23倍),相比带提示词重用的Megatron环形CP基线提高了2.63倍,相比不带重用的基线提高了7.08倍。

英文摘要

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.

发表机构

  • Fudan University(复旦大学)
  • The Pennsylvania State University(宾夕法尼亚州立大学)
  • Singapore Management University(新加坡管理大学)

机构由 AI 辅助整理,请以论文原文为准。

↑