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