Line data Source code
1 : /**
2 : * Copyright (c) 2025 Huawei Technologies Co., Ltd.
3 : * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4 : * CANN Open Software License Agreement Version 2.0 (the "License").
5 : * Please refer to the License for details. You may not use this file except in compliance with the License.
6 : * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7 : * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8 : * See LICENSE in the root of the software repository for the full text of the License.
9 : */
10 :
11 : #ifndef COLL_BATCH_SEND_RECV_GROUP_EXECUTOR_H
12 : #define COLL_BATCH_SEND_RECV_GROUP_EXECUTOR_H
13 :
14 : #include "coll_comm_executor.h"
15 : #include "coll_batch_send_recv_executor.h"
16 : #include <map>
17 :
18 : namespace hccl {
19 : class CollBatchSendRecvGroupExecutor : public CollBatchSendRecvExecutor {
20 : public:
21 : CollBatchSendRecvGroupExecutor(const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher);
22 128 : ~CollBatchSendRecvGroupExecutor() override = default;
23 : HcclResult Orchestrate(OpParam& param, AlgResourceResponse& algResource) override;
24 :
25 : protected:
26 : /* *************** 算法编排 *************** */
27 : u64 CalcSendLoopMaxCount(const u32 unitSize) const;
28 : u64 CalcRecvLoopMaxCount(const u32 unitSize) const;
29 : HcclResult CalcSendSlices();
30 : HcclResult CalcRecvSlices();
31 : HcclResult OrganizeSendItemByStream();
32 : HcclResult OrganizeRecvItemByStream();
33 : HcclResult CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport) override;
34 : struct SendRecvSlice {
35 : u8* addr;
36 : u64 size;
37 : u32 remoteRank;
38 : bool isRdma; // true=RDMA任务(用CCLOut拆半,无ping-pong);false=SDMA任务(用CCLIn,ping-pong)
39 37 : SendRecvSlice(u8* addr, u64 size, u32 remoteRank, bool isRdma = false)
40 37 : : addr(addr),
41 37 : size(size),
42 37 : remoteRank(remoteRank),
43 37 : isRdma(isRdma)
44 37 : {}
45 : };
46 :
47 : // SDMA slice按remoteRank % streamNum分发到各SDMA从流(仅SDMA,ping-pong)。
48 : std::vector<std::deque<SendRecvSlice>> sendDataSlicesBySendStream_;
49 : std::vector<std::deque<SendRecvSlice>> recvDataSlicesByRecvStream_;
50 :
51 : // RDMA与SDMA分离:SDMA占sendStreamNum_+recvStreamNum_条从流(CCLIn ping-pong);
52 : // RDMA单独占2条从流(1 send + 1 recv),对端顺序按alltoallv_direct_fullmesh对称场景规则
53 : // (send前向递增、recv后向递减,均跳过本pod;对端无任务则跳过)。
54 : // RDMA使用CCLOut拆半:A(offset 0)=send scratch(单slot), B(offset rdmaDataBlockSize_)=recv scratch(单slot)。
55 : // RDMA不做ping-pong,逐slice处理。
56 : static constexpr u32 RDMA_CCLOUT_HALF_NUM = 2; // CCLOut拆成两半:A=send, B=recv
57 : static constexpr u32 RDMA_STREAM_NUM = 2; // RDMA专用从流数:1 send + 1 recv
58 :
59 : // RDMA slice(已按对称场景规则排序),分别由rdmaSendStreamIdx_/rdmaRecvStreamIdx_对应的从流处理。
60 : std::deque<SendRecvSlice> rdmaSendSlices_;
61 : std::deque<SendRecvSlice> rdmaRecvSlices_;
62 :
63 : private:
64 : HcclResult RunLoop(OpParam& param);
65 : HcclResult RunTasks(OpParam& param);
66 : HcclResult ProcessPreloadedSendSlice(u32 streamIdx, u32& pendingSendCount, u32& nonEmptySendStream);
67 : HcclResult ProcessNewRankSendSlice(u32 streamIdx, u32& pendingSendCount, u32& nonEmptySendStream);
68 : HcclResult ProcessRecvSlice(u32 streamIdx, u32& nonEmptyRecvStream);
69 : // RDMA专用从流(单send/单recv)上处理一个slice,对端顺序由rdmaSendSlices_/rdmaRecvSlices_保证。
70 : HcclResult ProcessRdmaSendSlice();
71 : HcclResult ProcessRdmaRecvSlice();
72 : HcclResult CalcPodRange();
73 : bool IsRemoteRankRdma(u32 remoteRank) const;
74 : HcclResult SetNormalModeIfDeviceDirect();
75 : HcclResult CalcStreamNum(u32& streamNum) override;
76 : HcclResult CalcPingPongHalfSize();
77 :
78 : HcclResult MainPostSubWait(Stream& mainStream);
79 : HcclResult MainWaitSubPost(Stream& mainStream);
80 : // 统计各从流(SDMA send/recv + RDMA send/recv)是否有任务并记录到成员变量,
81 : // 同时返回SDMA非空流数供RunTasks循环使用。循环中不再更新。
82 : HcclResult CalcStreamTaskStatus(u32& nonEmptySendStream, u32& nonEmptyRecvStream);
83 :
84 : // 对称场景规则:send前向递增/recv后向递减遍历,跳过本pod。返回当前候选rank并推进游标。
85 : u32 GetNextDstRank(u32& curDstRank);
86 : u32 GetPreSrcRank(u32& curSrcRank);
87 : // 计算是否为非对称超节点场景(参考alltoallv_direct_fullmesh executor中的isSuPodAsym判断)
88 : void CalcIsSuPodAsym(bool isA2MultiModule);
89 : // 按对称场景规则对RDMA slice排序:isSend=true用前向(GetNextDstRank),false用后向(GetPreSrcRank)。
90 : // 遍历所有跨pod对端,仅输出存在任务的对端的slice(起点初始化与每次更新均跳过无任务对端)。
91 : void OrderRdmaSlices(
92 : bool isSend, const std::map<u32, std::deque<SendRecvSlice>>& byRank, std::deque<SendRecvSlice>& out);
93 :
94 : // RDMA专用从流在slaveStreams中的索引。
95 8 : u32 RdmaSendStreamIdx() const { return sendStreamNum_ + recvStreamNum_; }
96 4 : u32 RdmaRecvStreamIdx() const { return sendStreamNum_ + recvStreamNum_ + 1; }
97 :
98 : private:
99 : std::vector<std::deque<HcclSendRecvItem*>> sendQueueBySendstream_;
100 : std::vector<std::deque<HcclSendRecvItem*>> recvQueueByRecvstream_;
101 : u32 sendStreamNum_ = 0;
102 : u32 recvStreamNum_ = 0;
103 : u64 bufferSliceSize_ = 0;
104 : u64 rdmaDataBlockSize_ = 0; // RDMA单流slot大小 = CCLOut/2(A=send半区, B=recv半区)
105 : u32 podStartRank_ = 0; // pod(rank)范围,用于判定isRdma:[podStartRank_, podEndRank_]内为SDMA,跨pod为RDMA
106 : u32 podEndRank_ = 0;
107 : u32 devNumInlocalPod_ = 0; // 本pod内rank数,用于对称场景规则起点的计算
108 : bool isSuPodAsym_ = false; // 非对称场景:A2A3卡数不一致或A3多超节点server数不同时为true,send/recv使用相同遍历顺序
109 : // Ping-pong state
110 : std::vector<u32> sendCurPhase_; // which half has loaded data, ready to Record+Send (0=A, 1=B)
111 : std::vector<u64> sendLoadedSize_; // size of data loaded in current sendCurPhase_ half (0 = nothing)
112 : std::vector<u32> sendLoadedRemoteRank_; // remote rank for the data loaded in current sendCurPhase_
113 : std::vector<u32> recvCurPhase_; // which half to read from for recv (0=A, 1=B)
114 : std::vector<u32> recvCurRemoteRank_; // current remote rank being received on this stream
115 :
116 : // 各从流是否有任务(在RunTasks起始时记录,循环中不更新),用于头尾同步只唤醒有任务的从流。
117 : std::vector<bool> sendStreamHasTask_; // send从流(共sendStreamNum_条)是否有任务
118 : std::vector<bool> recvStreamHasTask_; // recv从流(共recvStreamNum_条)是否有任务
119 : bool rdmaSendHasTask_ = false; // RDMA专用send从流是否有任务
120 : bool rdmaRecvHasTask_ = false; // RDMA专用recv从流是否有任务
121 : };
122 : } // namespace hccl
123 :
124 : #endif
|