:从 CP2TP、Ring CP 到 True CP 的设计取舍与源码实现)
flash-linear-attention 上下文并行CP从 CP2TP、Ring CP 到 True CP 的设计取舍与源码实现【免费下载链接】flash-linear-attention Efficient implementations for emerging model architectures项目地址: https://gitcode.com/GitHub_Trending/fl/flash-linear-attentionflash-linear-attentionFLA为线性注意力 / 循环模型提供了一套上下文并行Context ParallelCP基础设施其设计取舍集中体现在 tests/context_parallel/README.md 中CP 并非只有一种形态仓库明确区分了 CP2TP头张量并行、Ring CP 与 True CP 三种方案并说明 FLA 内核中CP一词实际指的是 True CP 这一类真并行的分块状态合并方案。读完本文你将理解这三种方案的通信与依赖差异、build_cp_context的构建流程与参数约束以及all-gather merge数据流在 GDN/KDA 等 delta-rule 算子中的落地方式与测试验证方法。一、三种 CP 方案README 给出的核心界定tests/context_parallel/README.md 用三条界定把 CP 的三个常见方案区分开这也是理解整个上下文并行测试套件tests/context_parallel/的前提CP2TPContext Parallelism with Tensor Parallelism on attention heads按注意力头做张量并行。文档明确指出当在 FLA 内核中谈到 CP 时特指这种 TP头并行形态它需要两次 all-to-all 集合通信——一次在注意力计算之前一次在之后。也就是说它的并行本质是把每个 rank 的计算拆到不同头上而不是把序列维做真并行。Ring CP在 rank 之间引入顺序依赖。以 CP2 为例Rank1 必须等 Rank0 完成才能开始计算反向传播时反过来Rank0 又要等 Rank1。依赖环意味着各 rank 的启动时间是错峰的通信与计算难以完全重叠。True CP即 FLA 中实际实现的 CP文档强调KDA 与 GDN 最先支持。它让所有 rank 真正并行地计算通信开销最小但实现更复杂需要精巧的数学优化即后文的转移矩阵链式合并。这三条界定的实际意义在于FLA 仓库里cp_context、tests/context_parallel/下的test_cp_*.py全部针对的是第三种——True CP。它内部又名KCPKimi Context Parallel是 Moonshot AI 引入后独立贡献给 FLA 的实现见 fla/ops/cp/README.md 的致谢部分。二、True CP 的核心思想all-gather mergeTrue CP 要解决的问题是每个 rank 只持有序列的一个分片但循环状态依赖前面所有 token。FLA 的解法是把跨 rank 的状态依赖压缩成两个张量再借助集合通信一次性地重建初始状态。2.1 数据流局部计算 → all-gather → 链式合并按 fla/ops/cp/README.md 的 CP Architecture 一节局部计算每个 rank 用自己的本地 chunk 计算两个量S_ext ∈ R^{d_k × d_v}假设初始状态为零时累积出的状态M ∈ R^{d_k × d_k}转移矩阵刻画该 chunk 如何变换传入状态。All-gather收集所有 rank 的[S_ext, M]。Mergerank r 用前序 rank 的(M_j, S_ext,j)链式重建自己的初始状态$$\mathbf{S} \mathbf{0}; \quad \text{for } j (r - n_\text{pre}) \text{ to } (r-1): \quad \mathbf{S} \leftarrow \mathbf{M}j \mathbf{S} \mathbf{S}{\text{ext},j}$$关键点在于由于M与S_ext都可以独立于上游状态预先算出所有 rank 的 pre-process 阶段完全并行唯一的全局同步是一次 all-gather——这正是 True CP 相对 Ring CP最小通信、无顺序依赖优势的来源。2.2 转移矩阵的数学形式GDN/KDA 都建立在 delta rule 之上跨 chunk 的状态递推为$$\mathbf{S}{[t1]} \mathrm{Diag}(\gamma^C{[t]}) \mathbf{S}{[t]} (\boldsymbol{\Gamma}^{i \to C}{[t]} \odot \mathbf{K}{[t]})^\top (\mathbf{U}{[t]} - \mathbf{W}{[t]} \mathbf{S}{[t]})$$这正是 WY 表示下的形式对传入状态的作用恰好是对角衰减 − 二次项于是每个子 chunk 的转移矩阵为$$\mathbf{M}{[t]} \mathrm{Diag}(\gamma^C{[t]}) - (\boldsymbol{\Gamma}^{i \to C}{[t]} \odot \mathbf{K}{[t]})^\top \mathbf{W}{[t]}, \qquad \mathbf{M} \leftarrow \mathbf{M}{[t]} \mathbf{M}$$反向传播的结构与正向相同但方向相反从当前 rank 之后合并以回传梯度且dM是M的转置形式W^T KvsK^T W。该文档还强调一条工程红线M的链式相乘必须保持 fp32——在 bf16 下反复把 fp32 累加器转回 bf16会让误差随 chunk 数显著放大。2.3 前处理 / 反向预处理的内核实现上述数学在 fla/ops/cp/chunk_delta_h.py 中由 Triton 内核落地。pre_process_fwd_kernel_merged在一个 kernel 里同时产出 hV 部分与 mK 部分并用启发式参数区分不同模型的门控路径USE_GGDN 的标量门g非空时启用内核内部执行S ← γ^C · S的标量衰减USE_GKKDA 的逐维门gk非空时启用衰减变为逐维的S ← Diag(γ^C) SUSE_BGDPLRRWKV-7路径w与bg共享k的头维。这与 fla/ops/cp/README.md Gate Handling 一节的设计原则一致GDN 的标量门在内核里处理成本低所以 pre-process 与主内核都传原始k、gKDA 的Diag(α_t)是d_k × d_k矩阵内核内处理昂贵因此在 WY 表示阶段预门控出kgk * exp2(gk_last - gk)与qgq * exp2(gk)pre-process 和主内核必须接收完全相同的张量——这是正确性约束fla/ops/cp/README.md 末尾的 Input Tensor Summary 表格逐一列出了pre_process_fwd/fwd_h/pre_process_bwd/bwd_dhu四个阶段 GDN 与 KDA 各自的张量映射。2.4 通信原语集合通信封装在 fla/ops/cp/comm.pyall_gather_into_tensor(inp, out, group, async_op)CP 的核心集合操作输出形状[world_size, *inp.shape]all_reduce_sum求和归约send_recv_fwd/send_recv_bwd面向 causal_conv1d 的 CP 路径。由于卷积只看前W-1个 tokenrank 需要从前一个 rank 拿到序列尾部heads、向下一个 rank 传自己的尾部tails实现上通过 all-gather 保证所有 rank 参与rank 0 的前驱、末 rank 的后继返回零张量。这些函数均通过 fla/ops/cp/init.py 导出send_recv_fwd的 send to next / recv from prev 语义与 Ring CP 的点对点形态不同——它刻意避免了 ring 式的顺序依赖这也呼应了第一节 README 对 Ring CP 的批评。三、构建 CP 上下文build_cp_context的参数与约束3.1 基本用法fla/ops/cp/README.md 给出的 Quick Start 以全局cu_seqlens为输入from fla.ops.cp import build_cp_context # 分区前的全局 cu_seqlensdevice 可为 CPU 或 GPU cu_seqlens_global torch.tensor( [0, s1, s1 s2, ..., total], dtypetorch.long, devicedevice ) # causal_conv1d 的 CP 路径需要 conv1d_kernel_size cp_context build_cp_context( cu_seqlens_global, groupdist.group.WORLD, conv1d_kernel_sizeW, )3.2FLACPContext的字段语义build_cp_context返回的数据类定义在 fla/ops/cp/context.pyFLACPContext第 22–84 行关键字段包括字段含义groupCP 通信用的ProcessGroupis_cp_enabled即group is not Nonecu_seqlens/cu_seqlens_cpurank 本地的变长元数据GPU int64 / CPU 副本由全局cu_seqlens自动切分得到pre_num_ranks/post_num_ranks前驱 / 后继 rank 数merge 方向的边界依据is_first_rank/is_last_rank是否链首 / 链尾conv1d_kernel_size/pre_num_conv_tokens卷积 CP 路径所需的核宽与前序 token 深度use_tf32x3_affine_chainaffine-chain 点积用 tf32x3 提升精度仅 NVIDIAlayoutcontiguous默认或zigzag分片布局其中layout值得展开contiguous把 rank r 的 token 区间定为[r·T/W, (r1)·T/W)zigzag把链切成2W个等长区间rank r 持有第r段前部与第2W-1-r段后部调用方以[front; back]的拼接作为本地输入。zigzag 布局让因果注意力型的 CP 分片更均衡且无需在层边界重分片——zigzag 相关字段part_len、*_by_part元组等也都在该 dataclass 中逐一分列。3.3 源码中的硬性约束get_cp_cu_seqlensfla/ops/cp/context.py 第 128–207 行对输入做了明确校验total_tokenscu_seqlens[-1]必须不小于world_sizecontiguous或2·world_sizezigzag且必须能被其整除否则抛出ValueError变长序列通过searchsorted找出与本地区间相交的序列把全局cu_seqlens截断并平移为本地坐标int32存储该函数带tensor_cache装饰器相同输入不会重复计算分区元数据函数还计算pre_num_conv_tokens本地区间首序列在区间之前延伸的 token 数供 causal_conv1d 的 CP 路径取回卷积需要的上下文尾部。也就是说文档 Quick Start 中变长输入以全局cu_seqlens开始build_cp_context自动转换为 rank 本地元数据这句话在源码里就是上面这段 CPU 侧向量化切分逻辑。四、把算子接入 CPGDN 与 KDA 的调用链CP 模式下的算子调用遵循主内核前插入一次 pre-process的固定结构fla/ops/cp/README.md 的 Code Flow 一节给出了完整示例以 GDN 为例g chunk_local_cumsum(g, chunk_size64, scaleRCP_LN2) w, u recompute_w_u_fwd(k, v, beta, A, gg) # CP pre-process传入原始 k、标量 gUSE_GTrue, USE_GKFalse initial_state chunk_gated_delta_rule_fwd_h_pre_process( kk, ww, uu, gg, contextcp_context, ) # 主内核张量与 pre-process 完全一致 h, v_new, _ chunk_gated_delta_rule_fwd_h( kk, ww, uu, gg, initial_stateinitial_state, )KDA 的差异只在张量换成预门控版本kkg, gkg与qqg反向链路则为recompute_w_u_fwd → fwd_h → chunk_kda_bwd_dAv → bwd_dhu_pre_process → bwd_dhu。对外部使用者的约束fla/ops/cp/README.md 的 NOTECP 变长模式要求B 1rank 本地的cu_seqlens从 context 中取不要手动传cu_seqlens/cu_seqlens_cpuCP 模式下不支持initial_state与output_final_stateTrue——跨 rank 的状态同步由 merge 完成用户不提供初始状态causal_conv1d 的 CP 路径要求cp_context.conv1d_kernel_size与cp_context.cu_seqlens均已设置即build_cp_context时传入conv1d_kernel_size。五、测试与验证如何跑通并确认正确性5.1 测试矩阵tests/context_parallel/ 下的测试套件按算子组织test_cp_gdn.py、test_cp_gdn2.py、test_cp_kda.py、test_cp_gdp.py、test_cp_dplr.py、test_cp_rwkv7.py、test_cp_conv.py、test_cp_token_shift.py等并含test_cp_context.py分区元数据与test_cp_bwd_gk_offset.py等边界用例。以 tests/context_parallel/test_cp_gdn.py 为例其场景设计覆盖了 True CP 的典型切分形态测试配置考察点test_cp2_sequence_cutCP2, T10240,lengths[3000, 4000, 3240]序列在 rank 边界中间被截断test_cp2_boundary_alignedCP2,lengths[5120, 5120]序列边界与 rank 边界对齐test_cp4_complexCP4,lengths[7000, 3240]首序列横跨 3 个 ranktest_cp4_single_sequence/test_cp8_single_sequenceCP4/CP8 单长序列CP8 达 65536 token长程状态链式合并test_cp2_many_short_sequencesCP2, 7 条短序列多序列交替跨越 ranktest_cp2_gqa_*Hq HGQA 头映射Hq/Htest_cp2/4_state_v_firststate_v_firstTrue转置状态布局[V, K]验证方法是CP 路径 vs 单卡非 CP 参考的差分测试rank 0 先在整条序列上跑一遍chunk_gated_delta_rulevarlen、无 CP作为参考并取输出与dq/dk/dv/dg/db梯度随后各 rank 用build_cp_context构建上下文、切出本地分片跑 CP 前向与反向再all_gather拼接后逐项assert_close(ratio3e-3)。测试通过torch.multiprocessing.spawn启动多进程start_methodspawnNCCL 后端因此可以直接pytest tests/context_parallel/test_cp_gdn.py运行GPU 数量不足 2/4/8 时对应用例自动 skip文件末尾也保留了torchrun方式的手动入口setup_distributed_torchrun。5.2 调试方法论同目录下的 tests/context_parallel/debug.md 是一份 KCP 精度故障排查手册几个要点与 True CP 的实现细节直接对应按 token 比较不要按 chunk 比较KCP 与非 KCP 在同一变长序列上的 chunk 切分点不同该文档以lengths[400, 624]、chunk_size64、CP2 为例推演了两个 rank 的本地 chunk 边界逐 chunk 的h/dh在数学上不可比只有 per-token 张量v_new、bwd_dhu的中间dv、最终输入梯度才可比。只有当 KCP 切分点恰好落在每条序列的 chunk 边界上如均匀单序列时chunk 级比较才有意义。h/dh的语义约定两者都存储在 chunk起点进入该 chunk 的状态。KCP 中 rank r 的正向 merge 产出的是其首条本地序列的initial_state反向 merge 产出的是其末条本地序列的dht——对照非 KCP 时必须按全局 token 索引对齐不能信 chunk 下标。典型陷阱compress_h0压缩后的initial_state若在 autograd 中保存了原始入参KCP 下为Nonerank≥1 会在反向重算时静默丢失合并状态梯度整体偏离 3–5%expand_h0必须先于反向的前向重算执行merge_fwd_bwd_kernel的BV由 autotune 在{32, 64}中选择手写 grid 时若硬编码BV64会只填充一半 V 维表现为dh差异异常大。验收线变长 KCP、bf16、非对齐切分点下各梯度对 per-token naive 参考的norm_ratio应低于 5e-3纯 bf16 chunk-vs-per-token 噪声量级高于该值通常意味着反向重算用了错误的initial_state或 merge 内核拿到过期BV。工程提示调试时搭建单设备模拟所有 rank的 simulatorrun_nocp/run_cp/naive三层对照比反复mp.spawn快得多运行中不要清~/.triton/cache否则懒编译会与在跑内核竞争导致FileNotFoundError。六、适用范围与扩展性fla/ops/cp/README.md 的 Discussion 一节明确CP 机制不限于 delta-rule 递推——任何能写成chunk 间状态转移 转移矩阵M 累积状态S_ext分块形式的线性注意力都可以套用同一套build_cp_context all-gather merge 基础设施模型相关的只有两点(M, S_ext)如何由本地 chunk 算出、merge 内核如何跨 rank 链接它们。按文档口径当前已实现并验证的算子为GDN、GDP、GDN2、KDA 与 DPLRRWKV-7其中 GDP 目前要求提供gfla/ops/cp/README.md 的 Test References 一节给出的tests/context_parallel/test_cp_conv.py、test_cp_gdp.py、test_cp_kda.py、test_cp_gdn2.py即对应验证入口。回到 tests/context_parallel/README.md 的三条界定CP2TP 本质是头并行加两次 all-to-allRing CP 用顺序依赖换实现简单而 FLA 内核所指的 CP 是 True CP——用一次 all-gather 加 fp32 的M链式合并换取所有 rank 的真并行与最小通信开销。理解这条从方案分类到build_cp_context约束再到pre-process/merge 内核的链路是阅读、扩展或调试该仓库上下文并行代码的正确起点。【免费下载链接】flash-linear-attention Efficient implementations for emerging model architectures项目地址: https://gitcode.com/GitHub_Trending/fl/flash-linear-attention创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考