| reduce scatter support tensorlist.size != world_size Co-authored-by: limuan<liyijie16@huawei.com> # message auto-generated for no-merge-commit merge: !44439 merge reduce_scatter_v2.9.0 into v2.9.0 reduce scatter support tensorlist.size != world_size Created-by: limuan Commit-by: limuan Merged-by: ascend-robot Description: <!-- PR描述模板更新日期:20260203 --> # 【合入来源】 - [x] 需求 - [ ] 问题单 - [ ] issue/工单 - [ ] 重构优化 - [ ] 资料更新 # 【修改方案】 本提案修改 torch_npu 的 ProcessGroupHCCL::reduce_scatter,使其输入支持范围与 PyTorch 社区(NCCL)保持一致。 PyTorch 2.2 的 reduce_scatter 对输入张量列表有形状/计数约束(列表长度须等于 world_size、每张量 numel 须等于输出 numel 等),2.3 起去除了这些约束。pta(torch_npu)当前的 reduce_scatter 与 PyTorch 2.2 实现一致:输入是 tensor list,包含多个 tensor,数量与卡数一致。这与社区后续版本不一致,需要兼容输入只有一个 tensor 的场景等用例,和社区保持一致。 核心改动:在 reduce_scatter 的 same_size 分支新增展平函数 flatten_for_reduce_scatter,替代原先复用的 flatten_for_scatter_gather——移除"输入张量数须等于 world_size""每张量 numel 须等于输出 numel"两条校验,保留列表长度一致与 input/output 同设备校验,并新增非空检查。Python 层 torch.distributed.reduce_scatter 签名不变。) 各输入形态数值示例 统一 world_size=4、output=[4]、fp32、SUM,输入约定同 §2.1(rank r 展平后第 k 元 = r*10+k)。need = 16(输出 numel × world_size);全局位置 p 的归约值在 p < have 时为 60+4p,p ≥ have 时零填充。下文列出每个 rank 的输入与输出。 **① 单 tensor**(input_list = [[16]],have = 16 = need): rank0 IN : [0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15] OUT: [60,64,68,72] rank1 IN : [10,11,12,13,14,15,16,17,18,19,20,21,22,23,24,25] OUT: [76,80,84,88] rank2 IN : [20,21,22,23,24,25,26,27,28,29,30,31,32,33,34,35] OUT: [92,96,100,104] rank3 IN : [30,31,32,33,34,35,36,37,38,39,40,41,42,43,44,45] OUT: [108,112,116,120] **② 张量数 < 卡数**(input_list = [[4],[4]],have = 8 < 16,尾部两 rank 零填充): rank0 IN: rank1 IN: rank2 IN: rank3 IN: [0, 1, 2, 3] [10,11,12,13] [20,21,22,23] [30,31,32,33] [4, 5, 6, 7] [14,15,16,17] [24,25,26,27] [34,35,36,37] OUT: [60,64,68,72] OUT: [76,80,84,88] OUT: [0,0,0,0] OUT: [0,0,0,0] **③ 张量数 > 卡数**(input_list = [[4]]*5,have = 20 > 16,第 5 个张量整体忽略): rank0 IN: rank1 IN: rank2 IN: rank3 IN: [0, 1, 2, 3] [10,11,12,13] [20,21,22,23] [30,31,32,33] [4, 5, 6, 7] [14,15,16,17] [24,25,26,27] [34,35,36,37] [8, 9,10,11] [18,19,20,21] [28,29,30,31] [38,39,40,41] [12,13,14,15] [22,23,24,25] [32,33,34,35] [42,43,44,45] [16,17,18,19] [26,27,28,29] [36,37,38,39] [46,47,48,49] ← 忽略 OUT: [60,64,68,72] OUT: [76,80,84,88] OUT: [92,96,100,104] OUT: [108,112,116,120] **④ 张量数 = 卡数,每张量 > 输出**(input_list = [[5]]*4,have = 20 > 16,展平后尾部 4 元忽略): rank0 IN: rank1 IN: rank2 IN: rank3 IN: [0, 1, 2, 3, 4] [10,11,12,13,14] [20,21,22,23,24] [30,31,32,33,34] [5, 6, 7, 8, 9] [15,16,17,18,19] [25,26,27,28,29] [35,36,37,38,39] [10,11,12,13,14] [20,21,22,23,24] [30,31,32,33,34] [40,41,42,43,44] [15,16,17,18,19] [25,26,27,28,29] [35,36,37,38,39] [45,46,47,48,49] ← 后4元忽略 OUT: [60,64,68,72] OUT: [76,80,84,88] OUT: [92,96,100,104] OUT: [108,112,116,120] 展平后每 rank 共 20 元,前 16 元参与归约,尾部 4 元(第 4 个张量的后 4 元)忽略。 **⑤ 张量数 = 卡数,每张量 < 输出**(input_list = [[3]]*4,have = 12 < 16,尾部 rank 零填充): rank0 IN: rank1 IN: rank2 IN: rank3 IN: [0, 1, 2] [10,11,12] [20,21,22] [30,31,32] [3, 4, 5] [13,14,15] [23,24,25] [33,34,35] [6, 7, 8] [16,17,18] [26,27,28] [36,37,38] [9,10,11] [19,20,21] [29,30,31] [39,40,41] OUT: [60,64,68,72] OUT: [76,80,84,88] OUT: [92,96,100,104] OUT: [0,0,0,0] **⑥ 2D 及多维**(input_list = [[5,4]],have = 20 > 16,展平后尾部 4 元忽略,同 ③): rank0 IN (shape [5,4]): rank1 IN (shape [5,4]): rank2 IN (shape [5,4]): rank3 IN (shape [5,4]): [0, 1, 2, 3] [10,11,12,13] [20,21,22,23] [30,31,32,33] [4, 5, 6, 7] [14,15,16,17] [24,25,26,27] [34,35,36,37] [8, 9,10,11] [18,19,20,21] [28,29,30,31] [38,39,40,41] [12,13,14,15] [22,23,24,25] [32,33,34,35] [42,43,44,45] [16,17,18,19] [26,27,28,29] [36,37,38,39] [46,47,48,49] ← 忽略 OUT: [60,64,68,72] OUT: [76,80,84,88] OUT: [92,96,100,104] OUT: [108,112,116,120] # 【资料变更】 不涉及 # 【接口变更】 不涉及 # 【功能验证】 1、基础功能验证 2、dtype*op*tensorshape(1D,2D)*input_tensor_list.size(==,>,< world_size)*all_input_numbel(==, >, < )all_output_numbel reduce_scatter算子输出与gpu结果对比  3、性能验证,修改前后pta调用reduce_scatter,用时基本无变化  # 【CheckList】 > PR提交人对以下CheckList自检项进行全量自检,自检通过或不涉及,均修改 [ ] 为 [x] - [ ] 代码注释完备,正确记录错误日志 - [ ] 代码实现进行了返回值、空指针等校验 - [ ] PR标题正确使用类型标签,如:feat、fix、refactor、docs、test等 - [ ] PR持续集成流水线(CI)执行通过,代码检查无异常 See merge request: Ascend/pytorch!44439 | 5 天前 |