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 : #include "hcclCommDfx.h"
11 : #include "ccu_rep_context_v1.h"
12 : #include "task_info.h"
13 :
14 : namespace hccl {
15 :
16 : std::shared_mutex HcclCommDfx::baseLock_;
17 : std::mutex HcclCommDfx::taskIdMutex_;
18 : std::unordered_map<std::string,std::unordered_map<u64, u32> > HcclCommDfx::channelRemoteRankId_;
19 : std::unordered_map<u32, u32> HcclCommDfx::streamIdToTaskId_;
20 116 : HcclCommDfx::HcclCommDfx() {
21 116 : }
22 :
23 116 : HcclCommDfx::~HcclCommDfx() {
24 116 : setAddTaskCallback_ = nullptr;
25 116 : setAddDpuTaskCallback_ = nullptr;
26 116 : }
27 :
28 116 : HcclResult HcclCommDfx::Init(u32 deviceId, const std::string& comTag, u32 myRankId) {
29 116 : if (initializedFlag_) {
30 0 : return HCCL_SUCCESS;
31 : }
32 116 : HCCL_INFO("[%s]deviceId[%u], comTag[%s], myRankId[%u]", __func__, deviceId, comTag.c_str(), myRankId);
33 116 : deviceId_ = deviceId;
34 116 : commTag_ = comTag;
35 116 : myRankId_ = myRankId;
36 : // 1. 如果mirrorTaskManager_为空,则创建新的MirrorTaskManager
37 116 : if (!mirrorTaskManager_) {
38 116 : mirrorTaskManager_ = std::make_unique<Hccl::MirrorTaskManager>(deviceId_, &Hccl::GlobalMirrorTasks::Instance(), false);
39 : }
40 :
41 : // 2. 创建Profiling管理类
42 116 : EXCEPTION_CATCH(profiling_ = std::make_unique<HcclCommProfiling>(deviceId_, mirrorTaskManager_.get()), return HCCL_E_PTR);
43 116 : CHK_RET(profiling_->Init());
44 :
45 : // 3. 注册回调
46 232 : setAddTaskCallback_ = [this](u32 streamId, u32 taskId, const Hccl::TaskParam &taskParam, u64 handle) {
47 0 : return this->AddTaskInfoCallback(streamId, taskId, taskParam, handle);
48 116 : };
49 233 : setAddDpuTaskCallback_ = [this](const Hccl::TaskParam &taskParam, u64 handle) {
50 1 : return this->AddDpuTaskInfoCallback(taskParam, handle);
51 116 : };
52 116 : initializedFlag_ = true;
53 116 : return HCCL_SUCCESS; // 初始化成功返回成功码
54 : }
55 :
56 1 : HcclResult HcclCommDfx::IsOpBase(bool &isOpBase) {
57 1 : auto currDfxOpInfo = mirrorTaskManager_->GetCurrDfxOpInfo();
58 1 : CHK_SMART_PTR_NULL(currDfxOpInfo);
59 1 : isOpBase = currDfxOpInfo->op_.opMode == Hccl::OpMode::OPBASE;
60 1 : HCCL_INFO("[%s] IsOpBase: %d", __func__, isOpBase);
61 1 : return HCCL_SUCCESS;
62 1 : }
63 :
64 : // 回调注册实现
65 1 : void HcclCommDfx::AddTaskInfoCallbackLog(const Hccl::TaskParam &taskParam, const std::unordered_map<u64, u32> &handleMap) const
66 : {
67 1 : if (LIKELY(HcclCheckLogLevel(HCCL_LOG_INFO) == 0)) {
68 0 : return;
69 : }
70 2 : for (size_t i = 0; i < taskParam.ccuDetailInfo->size(); ++i) {
71 1 : const Hccl::CcuProfilingInfo &profInfo = (*taskParam.ccuDetailInfo)[i];
72 2 : for (int idx = 0; idx < hcomm::CCU_MAX_CHANNEL_NUM; idx++) {
73 2 : if (profInfo.channelId[idx] == hcomm::INVALID_VALUE_CHANNELID) {
74 1 : break;
75 : }
76 1 : auto handleIt = handleMap.find(profInfo.channelHandle[idx]);
77 1 : if (handleIt == handleMap.end()) {
78 0 : continue;
79 : }
80 1 : HCCL_INFO("[%s]idx[%u]: channelId[%u], remoteRankId[%u], channelHandle[0x%llx]",
81 : __func__, idx, profInfo.channelId[idx], handleIt->second, profInfo.channelHandle[idx]);
82 : }
83 : }
84 : }
85 :
86 5 : HcclResult HcclCommDfx::AddTaskInfoCallback(u32 streamId, u32 taskId, const Hccl::TaskParam &taskParam, u64 handle) {
87 5 : u32 remoteRankId = INVALID_UINT;
88 5 : if (handle != INVALID_U64) {
89 1 : CHK_RET(GetChannelRemoteRankId(commTag_, handle, remoteRankId));
90 : }
91 :
92 5 : Hccl::PrintTaskLog(streamId, taskId, taskParam, remoteRankId);
93 :
94 5 : if (taskParam.taskType == Hccl::TaskParamType::TASK_CCU && taskParam.ccuDetailInfo != nullptr) {
95 2 : std::shared_lock<std::shared_mutex> rwLock(baseLock_);
96 2 : auto commIt = channelRemoteRankId_.find(commTag_);
97 2 : if (commIt == channelRemoteRankId_.end()) {
98 1 : HCCL_ERROR("[%s] commTag:[%s] not found in CCU batch lookup", __func__, commTag_.c_str());
99 1 : return HCCL_E_PARA;
100 : }
101 1 : const auto& handleMap = commIt->second;
102 1 : AddTaskInfoCallbackLog(taskParam, handleMap);
103 2 : for (size_t i = 0; i < taskParam.ccuDetailInfo->size(); ++i) {
104 1 : Hccl::CcuProfilingInfo &profInfo = (*taskParam.ccuDetailInfo)[i];
105 2 : for (int idx = 0; idx < hcomm::CCU_MAX_CHANNEL_NUM; idx++) {
106 2 : if (profInfo.channelId[idx] == hcomm::INVALID_VALUE_CHANNELID) {
107 1 : break;
108 : }
109 1 : auto handleIt = handleMap.find(profInfo.channelHandle[idx]);
110 1 : if (handleIt == handleMap.end()) {
111 0 : HCCL_ERROR("[%s] Failed to get remote rank for channelHandle[0x%llx]",
112 : __func__, profInfo.channelHandle[idx]);
113 0 : return HCCL_E_PARA;
114 : }
115 1 : profInfo.remoteRankId[idx] = handleIt->second;
116 : }
117 : }
118 2 : }
119 8 : HcclResult ret = mirrorTaskManager_->AddTaskInfo(streamId, taskId,
120 8 : remoteRankId, taskParam, mirrorTaskManager_->GetCurrDfxOpInfo(), taskParam.isMaster);
121 4 : CHK_RET(ret);
122 4 : return HCCL_SUCCESS;
123 : }
124 :
125 3 : HcclResult HcclCommDfx::AddDpuTaskInfoCallback(const Hccl::TaskParam &taskParam, u64 handle) {
126 3 : u32 streamId = dpuStreamId_;
127 3 : u32 taskId = GetTaskId(streamId);
128 3 : Hccl::TaskParam localTaskParam = taskParam;
129 3 : localTaskParam.aicpuTaskId = aicpuTaskId_;
130 3 : localTaskParam.npuDevId = deviceId_;
131 3 : HCCL_INFO("[%s] streamId[%u], taskId[%u], aicpuTaskId[%u], npuDevId[%u].", __func__, streamId, taskId, localTaskParam.aicpuTaskId, localTaskParam.npuDevId);
132 6 : return AddTaskInfoCallback(streamId, taskId, localTaskParam, handle);
133 3 : }
134 :
135 0 : HcclResult HcclCommDfx::SetCurrDfxOpInfo(std::shared_ptr<Hccl::DfxOpInfo> dfxOpInfo)
136 : {
137 0 : profiling_->SetCurrDfxOpInfo(dfxOpInfo);
138 0 : return HCCL_SUCCESS;
139 : }
140 :
141 : // HcclCommDfx接口实现 - 修改为返回HcclResult类型
142 97 : HcclResult HcclCommDfx::ReportAllTasks(bool cachedReq) {
143 97 : EXCEPTION_CATCH(profiling_->ReportAllTasks(cachedReq), return HCCL_E_PTR);
144 97 : return HCCL_SUCCESS;
145 : }
146 :
147 0 : HcclResult HcclCommDfx::ReportOp(u64 beginTime, bool cachedReq, bool opbased) {
148 0 : EXCEPTION_CATCH(profiling_->ReportOp(beginTime, cachedReq, opbased), return HCCL_E_PTR);
149 0 : return HCCL_SUCCESS;
150 : }
151 :
152 : // 返回值Mc2要改
153 5 : void HcclCommDfx::ReportMc2CommInfo(const Mc2CommInfo& mc2CommInfo) {
154 5 : profiling_->ReportMc2CommInfo(mc2CommInfo);
155 5 : }
156 :
157 0 : HcclResult HcclCommDfx::UpdateProfStat() {
158 0 : profiling_->UpdateProfStat();
159 0 : return HCCL_SUCCESS;
160 : }
161 :
162 0 : Hccl::MirrorTaskManager* HcclCommDfx::GetMirrorTaskManager() const {
163 0 : return mirrorTaskManager_.get();
164 : }
165 :
166 : // 将remoteRankId添加到channelRemoteRankId_表中
167 5 : void HcclCommDfx::AddChannelRemoteRankId(const std::string& commTag, u64 handle, u32 remoteRankId) {
168 5 : std::unique_lock<std::shared_mutex> rwLock(baseLock_);
169 5 : HCCL_INFO("[HcclCommDfx][AddChannelRemoteRankId] commTag:[%s], handle:[%lu], remoteRankId:[%u]", commTag.c_str(), handle, remoteRankId);
170 5 : channelRemoteRankId_[commTag][handle] = remoteRankId;
171 5 : }
172 :
173 : // 在channelRemoteRankId_表中对remoteRankId进行查找(原有逻辑补充返回值)
174 4 : HcclResult HcclCommDfx::GetChannelRemoteRankId(const std::string& commTag, u64 handle, u32& remoteRankId) {
175 4 : std::shared_lock<std::shared_mutex> rwLock(baseLock_);
176 4 : auto commIt = channelRemoteRankId_.find(commTag);
177 4 : if (commIt == channelRemoteRankId_.end()) {
178 1 : HCCL_ERROR("[HcclCommDfx]commTag:[%s] not found", commTag.c_str());
179 1 : return HCCL_E_PARA;
180 : }
181 3 : auto handleIt = commIt->second.find(handle);
182 3 : if (handleIt == commIt->second.end()) {
183 1 : HCCL_ERROR("[HcclCommDfx]handle not found,commTag:[%s],handle:[%lu]", commTag.c_str(), handle);
184 1 : return HCCL_E_PARA;
185 : }
186 2 : remoteRankId = handleIt->second;
187 2 : return HCCL_SUCCESS; // 查找成功补充返回成功码
188 4 : }
189 :
190 5 : HcclResult HcclCommDfx::ReportKernel(uint64_t beginTime, const std::string& commTag, const std::string& kernelName, uint32_t threadId, bool cachedReq) {
191 5 : CHK_RET(profiling_->ReportKernel(beginTime, commTag, kernelName, threadId, cachedReq));
192 5 : return HCCL_SUCCESS;
193 : }
194 :
195 65548 : u32 HcclCommDfx::GetTaskId(u32 streamId)
196 : {
197 65548 : std::lock_guard<std::mutex> lock(taskIdMutex_);
198 65548 : auto& taskIdRef = streamIdToTaskId_[streamId];
199 65548 : constexpr u32 TASK_ID_MODULO = 65536;
200 65548 : taskIdRef = (taskIdRef + 1) % TASK_ID_MODULO;
201 65548 : u32 retTaskId = taskIdRef;
202 65548 : return retTaskId;
203 65548 : }
204 :
205 3 : void HcclCommDfx::SetDpuStreamId(u32 dpuStreamId) {
206 3 : dpuStreamId_ = dpuStreamId;
207 3 : }
208 :
209 : }
|