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 867 : AclgraphCallback &AclgraphCallback::GetInstance()
32 : {
33 867 : static AclgraphCallback aclgraphCallback;
34 867 : 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 不再访问这些 tag
58 1 : HcclResult aicpuRet = commIt.first->AicpuKfcClearOpResLaunch(commIt.second);
59 1 : if (aicpuRet != HCCL_SUCCESS) {
60 0 : HCCL_RUN_WARNING("[%s] modelID[%llu] aicpu batch sync fail, tagCount[%zu] ret[%d]; "
61 : "skip host link surgery this batch, tagsRequiringHostCleanup_ entries retained",
62 : __func__, modelId, commIt.second.size(), aicpuRet);
63 0 : isResourceReleaseFailed = true;
64 : }
65 :
66 : // 2. host 端逐 tag 清自己 resMap_/tagStreamInfo_ 等 host 进程内状态,与 aicpu sync 独立;aicpu 失败也要清,否则 resMap_ 残留
67 2 : for (auto &newTag : commIt.second) {
68 1 : ret = commIt.first->ClearOpResource(newTag, true);
69 1 : if (ret != HCCL_SUCCESS) {
70 0 : HCCL_ERROR("[%s] modelID[%llu] tag[%s] host resource release fail, ret[%d]",
71 : __func__, modelId, newTag.c_str(), ret);
72 0 : isResourceReleaseFailed = true;
73 : }
74 1 : HCCL_DEBUG("[%s] modelID[%llu] tag[%s] host resource release finish", __func__, modelId, newTag.c_str());
75 : }
76 : // 3. aicpu 已 sync,host 端整批 ListCommonRemove + 三容器 erase race-free;aicpu 失败时跳过,tag 保留至析构
77 1 : if (aicpuRet == HCCL_SUCCESS) {
78 1 : (void)commIt.first->ClearAclgraphHostLinks(commIt.second);
79 : }
80 : }
81 :
82 1 : captureResMap_.erase(modelId);
83 1 : captureCallbackParamMap_.erase(modelId);
84 1 : if (isResourceReleaseFailed) {
85 0 : HCCL_RUN_WARNING("[%s] modelID[%llu] resource release partially failed", __func__, modelId);
86 : } else {
87 1 : HCCL_INFO("[%s] modelID[%llu] resource release success", __func__, modelId);
88 : }
89 :
90 1 : return isResourceReleaseFailed ? HCCL_E_INTERNAL : HCCL_SUCCESS;
91 3 : }
92 :
93 805 : void AclgraphCallback::CleanCaptureRes(HcclCommunicator *communicator)
94 : {
95 805 : if (communicator == nullptr) {
96 0 : return;
97 : }
98 :
99 805 : std::lock_guard<std::mutex> lock(resMutex_);
100 806 : for (auto &modelIt : captureResMap_) {
101 1 : if (modelIt.second.find(communicator) != modelIt.second.end()) {
102 1 : modelIt.second.erase(communicator);
103 : }
104 : }
105 :
106 805 : HCCL_INFO("[%s] communicator[%p] resource release success", __func__, communicator);
107 805 : }
108 :
109 : // 记录aclgraph下发的所有tag, 首次记录时注册aclgraph销毁回调
110 7 : HcclResult AclgraphCallback::InsertNewTagToCaptureResMap(HcclCommunicator *communicator,
111 : const std::string &newTag, const OpParam &opParam)
112 : {
113 7 : CHK_PTR_NULL(communicator);
114 6 : aclmdlRI rtModel = nullptr;
115 6 : bool isCapture = false;
116 6 : u64 modelId = 0;
117 :
118 6 : CHK_RET(GetStreamCaptureInfo(opParam.stream.ptr(), rtModel, isCapture));
119 6 : CHK_PTR_NULL(rtModel);
120 6 : CHK_RET(GetModelId(rtModel, modelId));
121 :
122 6 : std::lock_guard<std::mutex> lock(resMutex_);
123 6 : if (captureResMap_.find(modelId) == captureResMap_.end()) {
124 4 : captureCallbackParamMap_[modelId].modelId = modelId;
125 8 : aclError aclRet = aclmdlRIDestroyRegisterCallback(rtModel, AclgraphDestroyCallback,
126 4 : static_cast<void *>(&captureCallbackParamMap_[modelId]));
127 4 : CHK_PRT_RET(aclRet != ACL_SUCCESS, HCCL_ERROR("[%s] aclmdlRIDestroyRegisterCallback fail, modelId[%llu]",
128 : __func__, modelId), HCCL_E_RUNTIME);
129 4 : HCCL_INFO("[%s] aclmdlRIDestroyRegisterCallback success modelID[%llu]", __func__, modelId);
130 : }
131 6 : captureResMap_[modelId][communicator].insert(newTag);
132 6 : HCCL_DEBUG("[%s] captureResMap insert tag[%s] to modelID[%llu]", __func__, newTag.c_str(), modelId);
133 :
134 6 : return HCCL_SUCCESS;
135 6 : }
136 : } // namespace hccl
|