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 : #include "aclgraph_callback.h"
12 : #include "stream_utils.h"
13 :
14 : namespace hccl {
15 :
16 0 : void AclgraphDestroyCallback(void* fnData)
17 : {
18 0 : AclgraphDestroyCallbackParam* callbackParam = static_cast<AclgraphDestroyCallbackParam*>(fnData);
19 0 : if (callbackParam == nullptr) {
20 0 : HCCL_ERROR("[%s] callbackParam ptr is NULL", __func__);
21 0 : return;
22 : }
23 :
24 0 : HCCL_INFO("[%s] Entry modelID[%llu] CleanCaptureRes", __func__, callbackParam->modelId);
25 0 : HcclResult ret = AclgraphCallback::GetInstance().CleanCaptureRes(callbackParam->modelId);
26 0 : if (ret != HCCL_SUCCESS) {
27 0 : HCCL_ERROR("[%s] modelID[%llu] CleanCaptureRes failed", __func__, callbackParam->modelId);
28 : }
29 : }
30 :
31 872 : AclgraphCallback& AclgraphCallback::GetInstance()
32 : {
33 872 : static AclgraphCallback aclgraphCallback;
34 872 : return aclgraphCallback;
35 : }
36 :
37 12 : AclgraphCallback::~AclgraphCallback()
38 : {
39 12 : std::lock_guard<std::mutex> lock(resMutex_);
40 12 : captureResMap_.clear();
41 12 : captureCallbackParamMap_.clear();
42 12 : }
43 :
44 3 : HcclResult AclgraphCallback::CleanCaptureRes(u64 modelId)
45 : {
46 : HcclResult ret;
47 :
48 3 : std::lock_guard<std::mutex> lock(resMutex_);
49 3 : auto modelIt = captureResMap_.find(modelId);
50 3 : if (modelIt == captureResMap_.end()) {
51 2 : HCCL_ERROR("[%s] modelID[%llu] is not record", __func__, modelId);
52 2 : return HCCL_E_NOT_FOUND;
53 : }
54 :
55 1 : bool isResourceReleaseFailed = false;
56 2 : for (auto& commIt : modelIt->second) {
57 : // 1. 整批 RPC sync aicpu 端:一次 launch erase 所有 tag 的 7 map entry,内含 sync,返回后 aicpu 不再访问这些
58 : // tag
59 1 : HcclResult aicpuRet = commIt.first->AicpuKfcClearOpResLaunch(commIt.second);
60 1 : if (aicpuRet != HCCL_SUCCESS) {
61 0 : HCCL_RUN_WARNING(
62 : "[%s] modelID[%llu] aicpu batch sync fail, tagCount[%zu] ret[%d]; "
63 : "skip host link surgery this batch, tagsRequiringHostCleanup_ entries retained",
64 : __func__, modelId, commIt.second.size(), aicpuRet);
65 0 : isResourceReleaseFailed = true;
66 : }
67 :
68 : // 2. host 端逐 tag 清自己 resMap_/tagStreamInfo_ 等 host 进程内状态,与 aicpu sync 独立;aicpu 失败也要清,否则
69 : // resMap_ 残留
70 2 : for (auto& newTag : commIt.second) {
71 1 : ret = commIt.first->ClearOpResource(newTag, true);
72 1 : if (ret != HCCL_SUCCESS) {
73 0 : HCCL_ERROR(
74 : "[%s] modelID[%llu] tag[%s] host resource release fail, ret[%d]", __func__, modelId, newTag.c_str(),
75 : ret);
76 0 : isResourceReleaseFailed = true;
77 : }
78 1 : HCCL_DEBUG("[%s] modelID[%llu] tag[%s] host resource release finish", __func__, modelId, newTag.c_str());
79 : }
80 : // 3. aicpu 已 sync,host 端整批 ListCommonRemove + 三容器 erase race-free;aicpu 失败时跳过,tag 保留至析构
81 1 : if (aicpuRet == HCCL_SUCCESS) {
82 1 : (void)commIt.first->ClearAclgraphHostLinks(commIt.second);
83 : }
84 : }
85 :
86 : // 清理通信域上的AIV modelId, aclgraph上modelId是可以复用的
87 2 : for (auto& commIt : modelIt->second) {
88 1 : commIt.first->EraseCaptureModelId(modelId);
89 : }
90 :
91 1 : captureResMap_.erase(modelId);
92 1 : captureCallbackParamMap_.erase(modelId);
93 1 : if (isResourceReleaseFailed) {
94 0 : HCCL_RUN_WARNING("[%s] modelID[%llu] resource release partially failed", __func__, modelId);
95 : } else {
96 1 : HCCL_INFO("[%s] modelID[%llu] resource release success", __func__, modelId);
97 : }
98 :
99 1 : return isResourceReleaseFailed ? HCCL_E_INTERNAL : HCCL_SUCCESS;
100 3 : }
101 :
102 808 : void AclgraphCallback::CleanCaptureRes(HcclCommunicator* communicator)
103 : {
104 808 : if (communicator == nullptr) {
105 0 : return;
106 : }
107 :
108 808 : std::lock_guard<std::mutex> lock(resMutex_);
109 920 : for (auto& modelIt : captureResMap_) {
110 112 : if (modelIt.second.find(communicator) != modelIt.second.end()) {
111 3 : modelIt.second.erase(communicator);
112 : }
113 : }
114 :
115 808 : HCCL_INFO("[%s] communicator[%p] resource release success", __func__, communicator);
116 808 : }
117 :
118 : // 记录aclgraph下发的所有tag, 首次记录时注册aclgraph销毁回调
119 7 : HcclResult AclgraphCallback::InsertNewTagToCaptureResMap(
120 : HcclCommunicator* communicator, const std::string& newTag, const OpParam& opParam)
121 : {
122 7 : CHK_PTR_NULL(communicator);
123 6 : aclmdlRI rtModel = nullptr;
124 6 : bool isCapture = false;
125 6 : u64 modelId = 0;
126 :
127 6 : CHK_RET(GetStreamCaptureInfo(opParam.stream.ptr(), rtModel, isCapture));
128 6 : CHK_PTR_NULL(rtModel);
129 6 : CHK_RET(GetModelId(rtModel, modelId));
130 :
131 6 : std::lock_guard<std::mutex> lock(resMutex_);
132 6 : CHK_RET(RegisterDestroyCallbackLocked(rtModel, modelId));
133 6 : captureResMap_[modelId][communicator].insert(newTag);
134 6 : HCCL_DEBUG("[%s] captureResMap insert tag[%s] to modelID[%llu]", __func__, newTag.c_str(), modelId);
135 :
136 6 : return HCCL_SUCCESS;
137 6 : }
138 :
139 2 : HcclResult AclgraphCallback::RegisterModelId(HcclCommunicator* communicator, aclmdlRI rtModel, u64 modelId)
140 : {
141 2 : CHK_PTR_NULL(communicator);
142 2 : CHK_PTR_NULL(rtModel);
143 :
144 2 : std::lock_guard<std::mutex> lock(resMutex_);
145 2 : CHK_RET(RegisterDestroyCallbackLocked(rtModel, modelId));
146 2 : captureResMap_[modelId][communicator];
147 2 : HCCL_DEBUG("[%s] register communicator[%p] to modelID[%llu]", __func__, communicator, modelId);
148 :
149 2 : return HCCL_SUCCESS;
150 2 : }
151 :
152 8 : HcclResult AclgraphCallback::RegisterDestroyCallbackLocked(aclmdlRI rtModel, u64 modelId)
153 : {
154 8 : if (captureResMap_.find(modelId) != captureResMap_.end()) {
155 2 : return HCCL_SUCCESS;
156 : }
157 6 : captureCallbackParamMap_[modelId].modelId = modelId;
158 12 : aclError aclRet = aclmdlRIDestroyRegisterCallback(
159 6 : rtModel, AclgraphDestroyCallback, static_cast<void*>(&captureCallbackParamMap_[modelId]));
160 6 : CHK_PRT_RET(
161 : aclRet != ACL_SUCCESS,
162 : HCCL_ERROR("[%s] aclmdlRIDestroyRegisterCallback fail, modelId[%llu]", __func__, modelId), HCCL_E_RUNTIME);
163 6 : HCCL_INFO("[%s] aclmdlRIDestroyRegisterCallback success modelID[%llu]", __func__, modelId);
164 :
165 6 : return HCCL_SUCCESS;
166 : }
167 : } // namespace hccl
|