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 __OP_UNFOLD_CACHE_ENTRY_H__
12 : #define __OP_UNFOLD_CACHE_ENTRY_H__
13 :
14 : #include <cstdint>
15 : #include <vector>
16 :
17 : #include "dispatcher_task_types.h" // LinkType
18 : #include "stream_pub.h"
19 :
20 : // 确认ptr应该为空
21 : #define CHK_PTR_NOTNULL(ptr) \
22 : do { \
23 : if (UNLIKELY((ptr) != nullptr)) { \
24 : HCCL_ERROR( \
25 : "[%s] errNo[0x%016llx] ptr[%s] is 0x%016llx (should be null), return HCCL_E_INTERNAL", __func__, \
26 : HCCL_ERROR_CODE(HCCL_E_INTERNAL), #ptr, (ptr)); \
27 : return HCCL_E_INTERNAL; \
28 : } \
29 : } while (0)
30 :
31 : // 确认ptrPtr不应该为空, 但*ptrPtr应该为空
32 : #define CHK_PTRPTR_NULL(ptrPtr) \
33 : do { \
34 : CHK_PTR_NULL(ptrPtr); \
35 : CHK_PTR_NOTNULL(*(ptrPtr)); \
36 : } while (0)
37 :
38 : namespace hccl {
39 :
40 : // 记录算子展开的输入/输出的内存范围
41 : // 注意: 内存的分配销毁由外部DeviceMem控制, 这里只是记录基地址和内存大小
42 : struct OpUnfoldMemRange {
43 : explicit OpUnfoldMemRange();
44 : explicit OpUnfoldMemRange(const uint64_t curBaseAddr, const uint64_t curMemSize);
45 : explicit OpUnfoldMemRange(const OpUnfoldMemRange& other);
46 : ~OpUnfoldMemRange();
47 :
48 : const OpUnfoldMemRange& operator=(const OpUnfoldMemRange& other); // 拷贝赋值操作符
49 :
50 : HcclResult GetEndAddr(uint64_t& endAddr) const; // 获取当前内存范围的end addr (exclusive)
51 : HcclResult InRange(const uint64_t addr, bool& isInRange) const;
52 :
53 : bool isValid;
54 : uint64_t baseAddr;
55 : uint64_t memSize;
56 : };
57 :
58 : struct RefreshAddrInfo {
59 : explicit RefreshAddrInfo();
60 : explicit RefreshAddrInfo(const uint32_t curRankId, const uint8_t curMemType);
61 : explicit RefreshAddrInfo(const RefreshAddrInfo& other);
62 : ~RefreshAddrInfo();
63 :
64 : const RefreshAddrInfo& operator=(const RefreshAddrInfo& other); // 拷贝赋值操作符
65 :
66 : static constexpr uint8_t INVALID_MEMTYPE = 0;
67 : static constexpr uint8_t USER_INPUT_MEMTYPE = 1;
68 : static constexpr uint8_t USER_OUTPUT_MEMTYPE = 2;
69 : static constexpr uint8_t HCCL_INPUT_MEMTYPE = 3; // 只用于alltoallv下的rank判断
70 :
71 : // 注意: 如果是alltoallv的PrepareIntraData, 则rankId表示当前send对应的remote rank, 即使dst memory为local hccl input
72 : // 参考OpUnfoldCacheEntry::UpdateRefreshAddrInfoForAlltoallv
73 : uint32_t rankId; // 默认情况下表示sqeAddr在rankId下对应memType的内存范围内
74 : uint8_t memType; // 0: invalid; 1: user input; 2: user output; 3: hccl input
75 : };
76 :
77 : using FlipInfo = std::pair<size_t, uint16_t>; // first: zero-taskid SQE idx; second: flipnum
78 : using RanksIdx = std::pair<std::vector<uint32_t>, uint32_t>; // first: ranks; second: idx
79 : using RankRflag = std::pair<uint32_t, bool>; // first: rank; second; recv flag (1: recv相关; 0: send相关)
80 :
81 : // 每个remote rank各有两个NotifyId/SignalAddr, 分别用于send/recv count对应的Wait/Record同步
82 : constexpr uint32_t NOTIFY_NUM_PER_REMOTE_RANK = 2;
83 :
84 : // 每个通信域只需要设置一次 (只由HCCL_BUFFSIZE和通信域拓扑决定, 与OpUnfoldCacheKey相关字段无关, e.g., opType and
85 : // workflowType)
86 : struct AlltoallvMetadata {
87 : // alltoallv第一次Orchestrate之前初始化
88 : uint64_t sdmaDataBlockSize = 0; // alltoallv的SDMA data block size (给定通信域下, 由于HCCL input buffer size,
89 : // SDMA并发数量, 以及deviceNumInLocalPod固定, 所以SDMA data block size也是固定的)
90 : std::vector<OpUnfoldMemRange>
91 : hcclInputMemRanges; // 每个rank的HCCL input buffer memory range (给定通信域, 在初始化后即固定)
92 : std::unordered_map<uint32_t, RankRflag>
93 : notifyIdRankRflagMap; // 跨卡通信的notifyId到remote RankRflag的映射 (用于NotifyWait的刷新)
94 : std::unordered_map<uint64_t, RankRflag>
95 : signalAddrRankRflagMap; // 跨卡通信的signalAddr到remote RankRflag的映射 (用于WriteRecord的刷新)
96 :
97 : // alltoallv第一次Orchestrate之后初始化
98 : // 注意: local/remote hccl offset只由local/target rank以及sdmaDataBlockSize决定
99 : // 注意: 一个hccl offset可能对应多个remote rank, 需要用RanksIdx追踪多个remote ranks以及当前需要使用的remote
100 : // rank的索引
101 : std::unordered_map<uint64_t, RanksIdx>
102 : hcclOffsetDstRanksIdxMap; // 当前rank的hccl input buffer中的local hccl offset到remote dst RanksIdx的映射
103 : // (用于PrepareIntraData)
104 :
105 : AlltoallvMetadata();
106 :
107 : void Clear();
108 : HcclResult Check(const bool afterFirstOrch) const;
109 : };
110 :
111 : // 每次alltoallv算子执行时更新
112 : struct AlltoallvSendRecvInfo {
113 : HcclDataType sendType = HcclDataType::HCCL_DATA_TYPE_RESERVED;
114 : HcclDataType recvType = HcclDataType::HCCL_DATA_TYPE_RESERVED;
115 : std::vector<uint64_t> sendCounts;
116 : std::vector<uint64_t> recvCounts;
117 : std::vector<uint64_t> sendOffsets;
118 : std::vector<uint64_t> recvOffsets;
119 :
120 : AlltoallvSendRecvInfo();
121 :
122 : HcclResult Check() const;
123 : };
124 :
125 : // 算子展开的动态缓存条目 (每个OpUnfoldKey对应最多一个缓存条目)
126 : class OpUnfoldCacheEntry {
127 : public:
128 : OpUnfoldCacheEntry() = delete;
129 : explicit OpUnfoldCacheEntry(
130 : const std::vector<OpUnfoldMemRange>& userInputMemRanges,
131 : const std::vector<OpUnfoldMemRange>& userOutputMemRanges);
132 : ~OpUnfoldCacheEntry();
133 :
134 : HcclResult GetSqeArrayCount(size_t& sqeArrayCount) const;
135 :
136 : // 缓存不命中下的函数
137 :
138 : // 分成两次函数调用是为了即使算子第一次展开的SQE存在placeholder, 一次LaunchTask下发的SQE仍然能够缓存在连续内存中,
139 : // 减少后续cache hit的开销
140 : HcclResult AllocSqeArray(
141 : const size_t sqeCount, const int32_t streamId,
142 : size_t& arrayIdx); // 分配成功会将arrayIdx设置为分配的SQE数组在sqeArrays_当中的索引
143 : HcclResult MemcpySqeArray(
144 : const size_t arrayIdx, const size_t sqeStartIdx, const size_t sqeCount, const uint8_t* sqeArray,
145 : const uint8_t* sqeTypeArray, const AicpuDfxInfo* sqeDfxInfoArray, const bool isAlltoallv,
146 : const AlltoallvMetadata*
147 : alltoallvMetadataPtr); // 将sqeArray memcpy到sqeArrays_[arrayIdx][sqeStartIdx:sqeStartIdx+sqeCount-1]
148 : // (因为DispatcherAicpu第一次算子展开时持有的是AlltoallvMetadata的指针,
149 : // 并且如果不是alltoallv算子则值为nullptr, 所以不传入引用)
150 :
151 : // 根据streamId计算streamSeqIdx
152 : HcclResult CalcStreamSeqIdxes(Stream& mainStream, std::vector<Stream>& slaveStreams);
153 :
154 : // 针对alltoallv类算子, 更新src/dst RefreshAddrInfo用于后续算子执行时的地址更新
155 : // (i) 更新invalid memType (只有cache-memcpy placeholder才可能出现此问题)
156 : // 当rankSize最后一个或多个ranks的send/recv count为0时, local user input/output offset为对应内存范围的end addr
157 : // -> 对于LocalCopy, src/dst memType默认为invalid, 需要更新为local user input/output
158 : // -> 对于PrepareIntraData, src memType默认为invalid, 需要更新为local user input
159 : // -> 对于RemoteCopy, dst memType默认为invalid, 需要更新为local user output
160 : // (ii) 更新local dst rank (如果dst memType是local hccl input)
161 : // PrepareIntraData场景下, 目的地址为local hccl offset, 因此dstRefreshInfo.rankId默认为local rank, 需要更新为remote
162 : // rank
163 : HcclResult UpdateRefreshAddrInfoForAlltoallv(const uint32_t curRank, AlltoallvMetadata& alltoallvMetadata);
164 :
165 : // 缓存命中下的函数
166 :
167 : // 更新指定的一段连续SQE, 并将相关信息设置给对应指针, 用于后续下发task到RTSQ
168 : // flipSqeIdxes指的是该段连续SQE中taskid==0且flipnum!=0的SQE的索引, 即这些SQE前面需要增加FlipPlaceholder
169 : HcclResult UpdateAndGetSqeArray(
170 : const size_t arrayIdx, const std::vector<OpUnfoldMemRange>& curUserInputMemRanges,
171 : const std::vector<OpUnfoldMemRange>& curUserOutputMemRanges, Stream& mainStream,
172 : std::vector<Stream>& slaveStreams, const uint32_t opRingBufferIdx, size_t& sqeCount, uint8_t** sqeArrayPtr,
173 : uint8_t** sqeTypeArrayPtr, AicpuDfxInfo** sqeDfxInfoArrayPtr, Stream** streamPtrPtr,
174 : std::vector<FlipInfo>& flipInfos, const bool profL1Enable, std::vector<uint64_t>& profTimestamps,
175 : const bool isAlltoallv, const AlltoallvMetadata& alltoallvMetadata,
176 : const AlltoallvSendRecvInfo& alltoallvSendRecvInfo);
177 :
178 : // Cache hit更新并下发entry中所有的SQE后, 由于缓存的SQE的addr-related fields被in-place更新,
179 : // 需要把userInputMemRanges_/userOutputMemRanges_为当前执行对应的memory ranges
180 : HcclResult SetInputOutputMemRanges(
181 : const std::vector<OpUnfoldMemRange>& curUserInputMemRanges,
182 : const std::vector<OpUnfoldMemRange>& curUserOutputMemRanges);
183 :
184 : private:
185 : // 合并两个uint32_t成为一个uint64_t
186 0 : inline void CombineUint32ToUint64(uint64_t& addr, const uint32_t high, const uint32_t low) const
187 : {
188 0 : constexpr uint64_t uintBitWidth = 32;
189 0 : addr = (static_cast<uint64_t>(high) << uintBitWidth) | static_cast<uint64_t>(low);
190 0 : return;
191 : }
192 :
193 : // 拆分一个uint64_t成为两个uint32_t
194 0 : inline void SplitUint64ToUint32(const uint64_t addr, uint32_t& high, uint32_t& low) const
195 : {
196 0 : constexpr uint64_t uintBitWidth = 32;
197 0 : high = static_cast<uint32_t>(addr >> uintBitWidth);
198 0 : low = static_cast<uint32_t>(addr & 0xFFFFFFFFULL);
199 0 : return;
200 : }
201 :
202 : // 缓存不命中下的函数
203 : HcclResult CheckAndPrepareRefreshAddrInfo(
204 : const uint64_t sqeAddr, RefreshAddrInfo& refreshAddrInfo, const bool isAlltoallv,
205 : const AlltoallvMetadata*
206 : alltoallvMetadataPtr); // 根据range判断sqeAddr是否在某个rankid的input/output user memory范围内,
207 : // 并相应更新RefreshAddrInfo为后续缓存命中刷新地址做准备
208 : // (因为DispatcherAicpu第一次算子展开时持有的是AlltoallvMetadata的指针,
209 : // 并且如果不是alltoallv算子则值为nullptr, 所以不传入引用)
210 : HcclResult CheckMemTypeForAlltoallv(
211 : const uint8_t* sqePtr, const uint8_t sqeType, const RefreshAddrInfo& srcRefreshAddrInfo,
212 : const RefreshAddrInfo& dstRefreshAddrInfo) const;
213 :
214 : // 缓存命中下的函数 (用于数据拷贝类SQE的刷新)
215 : HcclResult UpdateTransferSqeForAlltoallv(
216 : uint8_t* sqePtr, uint8_t* sqeTypePtr, const uint16_t curTaskId, const RefreshAddrInfo& srcRefreshAddrInfo,
217 : const RefreshAddrInfo& dstRefreshAddrInfo, const std::vector<OpUnfoldMemRange>& curUserInputMemRanges,
218 : const std::vector<OpUnfoldMemRange>& curUserOutputMemRanges, const AlltoallvMetadata& alltoallvMetadata,
219 : const AlltoallvSendRecvInfo&
220 : alltoallvSendRecvInfo); // 针对alltoallv算子刷新数据拷贝类的SQE (memcpy / cache-memcpy placeholder)
221 : HcclResult GetTransferCountForAlltoallv(
222 : uint64_t& count, uint64_t& size, const RefreshAddrInfo& srcRefreshAddrInfo,
223 : const RefreshAddrInfo& dstRefreshAddrInfo, const AlltoallvMetadata& alltoallvMetadata,
224 : const AlltoallvSendRecvInfo& alltoallvSendRecvInfo) const; // 针对alltoallv算子, 根据地址确定rank及数据拷贝大小
225 : HcclResult UpdateMemcpySqeForAlltoallv(
226 : uint8_t* sqePtr, uint8_t* sqeTypePtr, const uint16_t curTaskId, const RefreshAddrInfo& srcRefreshAddrInfo,
227 : const RefreshAddrInfo& dstRefreshAddrInfo, const std::vector<OpUnfoldMemRange>& curUserInputMemRanges,
228 : const std::vector<OpUnfoldMemRange>& curUserOutputMemRanges, const AlltoallvMetadata& alltoallvMetadata,
229 : const AlltoallvSendRecvInfo& alltoallvSendRecvInfo, const uint64_t count,
230 : const uint64_t size); // 针对alltoallv算子刷新Memcpy SQE
231 : HcclResult UpdateMemcpyPlaceholderSqeForAlltoallv(
232 : uint8_t* sqePtr, uint8_t* sqeTypePtr, const uint16_t curTaskId, const RefreshAddrInfo& srcRefreshAddrInfo,
233 : const RefreshAddrInfo& dstRefreshAddrInfo, const std::vector<OpUnfoldMemRange>& curUserInputMemRanges,
234 : const std::vector<OpUnfoldMemRange>& curUserOutputMemRanges, const AlltoallvMetadata& alltoallvMetadata,
235 : const AlltoallvSendRecvInfo& alltoallvSendRecvInfo, const uint64_t count,
236 : const uint64_t size); // 针对alltoallv算子刷新CacheMemcpyPlaceholder SQE
237 : HcclResult RefreshSqeAddr(
238 : uint64_t& sqeAddr, const uint32_t rankId, const std::vector<OpUnfoldMemRange>& cachedMemRanges,
239 : const std::vector<OpUnfoldMemRange>& curMemRanges, const bool isAlltoallv,
240 : const uint64_t offset) const; // 根据range判断是否需要刷新, 根据计算/给定的offset进行刷新
241 :
242 : // 缓存命中下的函数 (用于同步类SQE的刷新)
243 : HcclResult UpdateSyncSqeForAlltoallv(
244 : uint8_t* sqePtr, uint8_t* sqeTypePtr, const uint16_t curTaskId, const RefreshAddrInfo& srcRefreshAddrInfo,
245 : const RefreshAddrInfo& dstRefreshAddrInfo, const AlltoallvMetadata& alltoallvMetadata,
246 : const AlltoallvSendRecvInfo& alltoallvSendRecvInfo); // 针对alltoallv算子刷新同步类的SQE (notify / write-value /
247 : // cache-notify / cache-write)
248 : HcclResult GetTransferCountForAlltoallv(
249 : uint64_t& count, uint64_t& size, const uint8_t* sqePtr, const uint8_t* sqeTypePtr,
250 : const AlltoallvMetadata& alltoallvMetadata, const AlltoallvSendRecvInfo& alltoallvSendRecvInfo)
251 : const; // 针对alltoallv算子, 根据notifyId/signalAddr确定rank及数据拷贝大小
252 : HcclResult UpdateNotifyPlaceholderSqeForAlltoallv(
253 : uint8_t* sqePtr, uint8_t* sqeTypePtr, const uint16_t curTaskId, const AlltoallvMetadata& alltoallvMetadata,
254 : const AlltoallvSendRecvInfo& alltoallvSendRecvInfo, const uint64_t count,
255 : const uint64_t size); // 针对alltoallv算子刷新cache-notify placeholder
256 : HcclResult UpdateWritePlaceholderSqeForAlltoallv(
257 : uint8_t* sqePtr, uint8_t* sqeTypePtr, const uint16_t curTaskId, const AlltoallvMetadata& alltoallvMetadata,
258 : const AlltoallvSendRecvInfo& alltoallvSendRecvInfo, const uint64_t count,
259 : const uint64_t size); // 针对alltoallv算子刷新cache-write placeholder
260 : HcclResult UpdateNotifySqeForAlltoallv(
261 : uint8_t* sqePtr, uint8_t* sqeTypePtr, const uint16_t curTaskId, const AlltoallvMetadata& alltoallvMetadata,
262 : const AlltoallvSendRecvInfo& alltoallvSendRecvInfo, const uint64_t count,
263 : const uint64_t size); // 针对alltoallv算子刷新notify SQE
264 : HcclResult UpdateWriteValueSqeForAlltoallv(
265 : uint8_t* sqePtr, uint8_t* sqeTypePtr, const uint16_t curTaskId, const AlltoallvMetadata& alltoallvMetadata,
266 : const AlltoallvSendRecvInfo& alltoallvSendRecvInfo, const uint64_t count,
267 : const uint64_t size); // 针对alltoallv算子刷新WriteValue SQE
268 : HcclResult UpdateMemcpyRecordSqeForAlltoallv(
269 : uint8_t* sqePtr, uint8_t* sqeTypePtr, const uint16_t curTaskId, const AlltoallvMetadata& alltoallvMetadata,
270 : const AlltoallvSendRecvInfo& alltoallvSendRecvInfo, const uint64_t count,
271 : const uint64_t size); // 针对alltoallv算子刷新MemcpyRecord SQE
272 : HcclResult UpdateMemcpyRecordPlaceholderSqeForAlltoallv(
273 : uint8_t* sqePtr, uint8_t* sqeTypePtr, const uint16_t curTaskId, const AlltoallvMetadata& alltoallvMetadata,
274 : const AlltoallvSendRecvInfo& alltoallvSendRecvInfo, const uint64_t count,
275 : const uint64_t size); // 针对alltoallv算子刷新cache-memcpy-record placeholder SQE
276 : void SetCachePlaceholderHeaderForAlltoallv(const uint16_t streamId, const uint16_t taskId, uint8_t* sqePtr);
277 :
278 : std::vector<uint8_t*>
279 : sqeArrays_; // 多段连续的SQE数组 (每段连续的SQE不超过HCCL_SQE_SIZE * HCCL_PER_LAUNCH_SQE_CNT bytes)
280 : std::vector<uint8_t*> sqeTypeArrays_; // 每段每个SQE的type
281 : std::vector<AicpuDfxInfo*> sqeDfxInfoArrays_; // 每段每个SQE的DfxInfo
282 : std::vector<int32_t> streamIds_; // 每段SQE对应的actual stream ID
283 : std::vector<uint32_t> streamSeqIdxes_; // 每段SQE对应的sequential stream index (sequential是指将mainStream +
284 : // slaveStreams顺序起来看, 0代表mainStream, 1代表slaveStreams[0])
285 : std::vector<std::vector<RefreshAddrInfo>> srcRefreshAddrInfoArrays_; // 每段每个SQE中dstAddr (if any)对应的刷新信息
286 : std::vector<std::vector<RefreshAddrInfo>> dstRefreshAddrInfoArrays_; // 每段每个SQE中dstAddr (if any)对应的刷新信息
287 :
288 : std::vector<OpUnfoldMemRange> userInputMemRanges_; // 当前通信域每个rank的user input memory range
289 : std::vector<OpUnfoldMemRange> userOutputMemRanges_; // 当前通信域每个rank的user output memory range
290 : };
291 :
292 : }; // namespace hccl
293 :
294 : #endif // __OP_UNFOLD_CACHE_ENTRY_H__
|