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 HCCL_DISPATCHER_AICPU_PUB_H
12 : #define HCCL_DISPATCHER_AICPU_PUB_H
13 :
14 : #include <vector>
15 : #include <functional>
16 : #include "sal_pub.h"
17 : #include "dispatcher_pub.h"
18 :
19 : #include "aicpu/aicpu_hccl_sqcq.h"
20 : #include "aicpu/aicpu_hccl_sqcqv1.h"
21 : #include "aicpu/aicpu_hccl_sqcqv2.h"
22 :
23 : #include "op_unfold_cache.h"
24 :
25 : namespace hccl {
26 : using AddOneNotifyWaitSqe = void(*)(uint16_t, uint16_t, u64, const uint8_t *, uint8_t *, const dfx::DfxTimeOutConfig &);
27 : using AddOneRecordSqe = void(*)(uint16_t, uint16_t, u64, const uint8_t *, uint8_t *);
28 : using AddOneWriteValueRecordSqe = void(*)(uint16_t, uint16_t, u64, const uint8_t *, uint8_t *);
29 : using AddOneMemcpySqe = void(*)(uint16_t, uint16_t, const void *, uint32_t, const aclDataType,
30 : aclrtReduceKind, const void *, uint32_t, uint32_t, uint32_t, u64, uint8_t, const uint8_t *, uint8_t *, uint32_t);
31 : using AddOneEventResetSqe = void(*)(uint16_t, int32_t, uint16_t, int64_t, int64_t,
32 : u64, const uint8_t *, uint8_t *);
33 : using AddOneEventRecordSqe = void(*)(uint16_t, int32_t, uint16_t, const uint8_t *, uint8_t *);
34 : using AddOneEventWaitSqe = void(*)(uint16_t, int32_t, uint16_t, const uint8_t *, uint8_t *);
35 : using AddOneRdmaDbSendSqe = void(*)(uint16_t, uint16_t, uint64_t, uint64_t, uint32_t, uint8_t, const uint8_t *, uint8_t *);
36 : using AddOnePlaceHolderSqe = void(*)(uint16_t, uint16_t, uint16_t, const uint8_t *, uint8_t *);
37 : using AddOneCacheMemcpyPlaceHolderSqe = void(*)(uint16_t, uint16_t, const void *, const void *, uint8_t, const uint8_t *,
38 : uint8_t *, uint32_t);
39 : using AddOneCacheNotifyWaitPlaceholderSqe = void(*)(uint16_t, uint16_t, u64, const uint8_t *, uint8_t *, const dfx::DfxTimeOutConfig &);
40 : using AddOneCacheNotifyRecordPlaceholderSqe = void(*)(uint16_t, uint16_t, u64, const uint8_t *, uint8_t *);
41 : using AddOneCacheWriteValuePlaceholderSqe = void(*)(uint16_t, uint16_t, u64, const uint8_t *, uint8_t *);
42 : using AddOneCacheMemcpyRecordPlaceholderSqe = void(*)(uint16_t, uint16_t, const void *, uint32_t, const aclDataType,
43 : aclrtReduceKind, const void *, uint32_t, uint32_t, uint32_t, u64, uint8_t, const uint8_t *, uint8_t *, uint32_t);
44 :
45 : class DispatcherAiCpu : public DispatcherPub {
46 : public:
47 : explicit DispatcherAiCpu(const u32 devPhyId);
48 : ~DispatcherAiCpu() override;
49 : HcclResult Init() override;
50 : HcclResult WaitValue(hccl::Stream &stream, u64 waitAddr, u64 valueAddr, bool reset) override;
51 : HcclResult WriteValue(hccl::Stream &stream, u64 writeAddr, u64 valueAddr) override;
52 : HcclResult SignalRecord(HcclRtNotify signal, hccl::Stream &stream, u32 userRank, u64 offset = INVALID_U64,
53 : s32 stage = INVALID_VALUE_STAGE, bool inchip = false, u64 signalAddr = INVALID_U64,
54 : u32 notifyId = INVALID_UINT) override;
55 : HcclResult SignalRecord(hccl::DeviceMem &dst, hccl::DeviceMem &src, hccl::Stream &stream,
56 : u32 remoteUserRank, hccl::LinkType inLinkType, u32 notifyId) override;
57 : HcclResult SignalWait(HcclRtNotify signal, hccl::Stream &stream, u32 userRank, u32 remoteUserRank,
58 : s32 stage = INVALID_VALUE_STAGE, bool inchip = false, u32 notifyId = INVALID_UINT,
59 : u32 timeOut = NOTIFY_INVALID_WAIT_TIME) override;
60 : HcclResult MemcpyAsync(hccl::DeviceMem &dst, const hccl::DeviceMem &src, hccl::Stream &stream,
61 : u32 remoteUserRank = INVALID_VALUE_RANKID, hccl::LinkType inLinkType = hccl::LinkType::LINK_ONCHIP) override;
62 : HcclResult ReduceAsync(const void *src, void *dst, u64 dataCount, const HcclDataType datatype, HcclReduceOp redOp,
63 : Stream &stream, HcclReduceType reduceType = HcclReduceType::HCCL_TBE_REDUCE) override;
64 : HcclResult InlineReduceAsync(const void *src, u64 dataCount, const HcclDataType datatype, HcclReduceOp redOp,
65 : Stream &stream, void *dst, u32 remoteUserRank = INVALID_VALUE_RANKID,
66 : hccl::LinkType inLinkType = hccl::LinkType::LINK_ONCHIP) override;
67 : HcclResult RdmaRecord(u32 dbindex, u64 dbinfo, const struct SendWr &wr, hccl::Stream &stream,
68 : RdmaType rdmaType, u32 userRank, u64 offset, u32 notifyId) override;
69 :
70 : HcclResult LaunchTasksEx(Stream &stream, std::vector<Stream> &subStreams) override;
71 : HcclResult LaunchAllTasks() override;
72 :
73 : HcclResult RdmaSend(u32 dbindex, u64 dbinfo, hccl::Stream &stream, RdmaTaskInfo &taskInfo) override;
74 : // 新增接口用于算子展开的动态缓存
75 : HcclResult ClearLaunchContext(); // 当前算子展开不需要使用动态缓存
76 : // 设置launch context, 在LaunchTask时用于算子展开动态缓存的admission (因为需要在DispatcherAicpu中暂存AlltoallvMetadata, 所以传入指针而不是引用)
77 : HcclResult SetLaunchContext(const OpUnfoldKey& key, OpUnfoldCache *cachePtr,
78 : const std::vector<OpUnfoldMemRange>& userInputMemRanges, const std::vector<OpUnfoldMemRange>& userOutputMemRanges,
79 : const bool isAlltoallv, const AlltoallvMetadata* alltoallvMetadataPtr);
80 : // 缓存命中时, 使用缓存中的SQE信息下发给RTSQ
81 : HcclResult LaunchNewTask(OpUnfoldCacheEntry *entryPtr, const std::vector<OpUnfoldMemRange>& userInputMemRanges,
82 : const std::vector<OpUnfoldMemRange>& userOutputMemRanges, Stream& mainStream, std::vector<Stream> &slaveStreams,
83 : const bool profL1Enable, const bool isAlltoallv, const AlltoallvMetadata& alltoallvMetadata, const AlltoallvSendRecvInfo& alltoallvSendRecvInfo);
84 :
85 : HcclResult LaunchTask(Stream &stream, bool isBlockLaunch);
86 : HcclResult TbeReduceAsync(const void *src1, const void *src2, u64 count, const HcclDataType datatype,
87 : HcclReduceOp redOp, Stream &stream, const void *dst);
88 : HcclResult AddRetryPreamble(Stream &stream) override;
89 : HcclResult StreamSync(Stream &stream) override;
90 :
91 11 : void SetOpExecStatusCallback(std::function<HcclResult()> checkOpExecStatusCallback)
92 : {
93 11 : checkOpExecStatusCallback_ = checkOpExecStatusCallback;
94 11 : return;
95 : }
96 :
97 0 : void SetOpRingBufferIdx(const u32 opRingBufferIdx)
98 : {
99 0 : opRingBufferIdx_ = opRingBufferIdx;
100 0 : HCCL_INFO("[DispatcherAiCpu][SetOpRingBufferIdx]DFX opRingBufferIdx: [%u]",
101 : opRingBufferIdx);
102 0 : return;
103 : }
104 :
105 11 : void SetSqeTimeOut(const u64 timeOut)
106 : {
107 11 : if (timeOut > notifyMaxWaitTime_) {
108 0 : dfxTimeOutConfig_.sqeTimeOutTimeOut = notifyMaxWaitTime_;
109 0 : HCCL_WARNING("[SetSqeTimeOut] timeOut[%llu] exceeds the maximum allowed value "
110 : "for notifyMaxWaitTime[%u].", timeOut, notifyMaxWaitTime_);
111 : } else {
112 11 : dfxTimeOutConfig_.sqeTimeOutTimeOut = timeOut;
113 : }
114 11 : HCCL_INFO("[DispatcherAiCpu][SetSqeTimeOut]DFX timeout config init successfully with details: [%s]",
115 : dfxTimeOutConfig_.ToString().c_str());
116 11 : return;
117 : }
118 :
119 : void GetSqeTimeOut(u64 &timeOut)
120 : {
121 : timeOut = dfxTimeOutConfig_.sqeWaitTimeOut;
122 : return;
123 : }
124 :
125 0 : HcclResult SetSqFullWaitTimeOut(u64 notifyWaitTime)
126 : {
127 0 : dfxTimeOutConfig_.sqFullWaitTimeOut = (notifyWaitTime == 0) ?
128 : notifyWaitTime : (notifyWaitTime + AICPU_RTSQ_TIMEOUT_INC);
129 0 : HCCL_INFO("[DispatcherAiCpu][SetSqFullWaitTimeOut]DFX timeout config with details: [%s]",
130 : dfxTimeOutConfig_.ToString().c_str());
131 0 : return HCCL_SUCCESS;
132 : }
133 0 : HcclResult SignalRecord(Stream &stream, u64 notifyId)
134 : {
135 0 : return SignalRecord(nullptr, stream, INVALID_VALUE_RANKID, INVALID_U64, INVALID_VALUE_STAGE, true,
136 0 : INVALID_U64, static_cast<u32>(notifyId));
137 : }
138 0 : HcclResult SignalWait(Stream &stream, u32 notifyId, u32 timeOut)
139 : {
140 0 : return SignalWait(nullptr, stream, INVALID_VALUE_RANKID, INVALID_VALUE_RANKID,
141 0 : INVALID_VALUE_STAGE, true, static_cast<u32>(notifyId), timeOut);
142 : }
143 : public:
144 : dfx::DfxTimeOutConfig dfxTimeOutConfig_ = {0};
145 : uint32_t opRingBufferIdx_ = 0;
146 : private:
147 : // 新增接口用于算子展开的动态缓存
148 : HcclResult WaitRtsq(Stream& stream, const size_t& sqeCount, const bool isBlockLaunch); // 等待RTSQ直到有sqeCount的SQE的空间 (与LaunchTask中相同的逻辑)
149 : HcclResult MemcpyRtsq(Stream& stream, const size_t sqeCount, const uint8_t *sqeArray, const uint8_t *sqeTypeArray, const AicpuDfxInfo *sqeDfxInfoArray, const bool profL1Enable, const std::vector<uint64_t>& profTimestamps, const size_t profTimestampStartIdx); // 将动态缓存中更新后的SQE的相关信息下发到RTSQ中
150 :
151 : HcclResult AddFlipTask(Stream &stream);
152 : HcclResult GetStreamSqeBufferAddr(hccl::Stream &stream, uint8_t *&sqeBufferAddr, uint8_t *&sqeTypeAddr,
153 : uint8_t *&sqeDfxInfoAddr, uint16_t &taskId);
154 : void SaveStreamInfo(hccl::Stream &stream);
155 : u64 CalcDbAddr(u32 dbindex);
156 : void InitTimeOutConfig();
157 444 : u32 GetMaxNotifyWaitTime()
158 : {
159 444 : return notifyMaxWaitTime_;
160 : }
161 :
162 : AddOneNotifyWaitSqe addOneNotifyWaitSqe_ = nullptr;
163 : AddOneRecordSqe addOneRecordSqe_ = nullptr;
164 : AddOneWriteValueRecordSqe addOneWriteValueRecordSqe_ = nullptr;
165 : AddOneMemcpySqe addOneMemcpySqe_ = nullptr;
166 : AddOneEventResetSqe addOneEventResetSqe_ = nullptr;
167 : AddOneEventRecordSqe addOneEventRecordSqe_ = nullptr;
168 : AddOneEventWaitSqe addOneEventWaitSqe_ = nullptr;
169 : AddOneRdmaDbSendSqe addOneRdmaDbSendSqe_ = nullptr;
170 : AddOnePlaceHolderSqe addOneFlipPlaceHolderSqe_ = nullptr;
171 : AddOneCacheMemcpyPlaceHolderSqe addOneCacheMemcpyPlaceHolderSqe_ = nullptr;
172 : AddOneCacheNotifyWaitPlaceholderSqe addOneCacheNotifyWaitPlaceholderSqe_ = nullptr;
173 : AddOneCacheNotifyRecordPlaceholderSqe addOneCacheNotifyRecordPlaceholderSqe_ = nullptr;
174 : AddOneCacheWriteValuePlaceholderSqe addOneCacheWriteValuePlaceholderSqe_ = nullptr;
175 : AddOneCacheMemcpyRecordPlaceholderSqe addOneCacheMemcpyRecordPlaceholderSqe_ = nullptr;
176 : std::function<HcclResult()> checkOpExecStatusCallback_ = nullptr;
177 :
178 : HcclAicpuDispatcherInfo aicpuInfo_;
179 :
180 : std::unordered_map<s32, Stream> streamMap_; // 保存下过task的stream
181 : u64 notifySize_ = 0;
182 :
183 : // Launch context用于算子展开的动态缓存
184 : // 注意: cachePtr_初始化为空, needAddSqe_初始化为false, 即暂无算子展开的动态缓存
185 : OpUnfoldKey key_; // 当前展开算子的标识符
186 : OpUnfoldCache *cachePtr_ = nullptr; // 算子展开的动态缓存
187 : std::vector<OpUnfoldMemRange> userInputMemRanges_; // 当前算子展开执行时, 通信域内各rank分配的user input memory range
188 : std::vector<OpUnfoldMemRange> userOutputMemRanges_; // 当前算子展开执行时, 通信域内各rank分配的user output memory range
189 : bool isAlltoallv_ = false;
190 : const AlltoallvMetadata* alltoallvMetadataPtr_ = nullptr; // alltoallv算子对应的metadata (与通信域绑定)
191 : bool needAddSqe_ = false;
192 : };
193 : } // namespace hccl
194 : #endif // HCCL_DISPATCHER_AICPU_PUB_H
|