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 870 : AclgraphCallback& AclgraphCallback::GetInstance()
32 : {
33 870 : static AclgraphCallback aclgraphCallback;
34 870 : 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 1 : captureResMap_.erase(modelId);
87 1 : captureCallbackParamMap_.erase(modelId);
88 1 : if (isResourceReleaseFailed) {
89 0 : HCCL_RUN_WARNING("[%s] modelID[%llu] resource release partially failed", __func__, modelId);
90 : } else {
91 1 : HCCL_INFO("[%s] modelID[%llu] resource release success", __func__, modelId);
92 : }
93 :
94 1 : return isResourceReleaseFailed ? HCCL_E_INTERNAL : HCCL_SUCCESS;
95 3 : }
96 :
97 808 : void AclgraphCallback::CleanCaptureRes(HcclCommunicator* communicator)
98 : {
99 808 : if (communicator == nullptr) {
100 0 : return;
101 : }
102 :
103 808 : std::lock_guard<std::mutex> lock(resMutex_);
104 809 : for (auto& modelIt : captureResMap_) {
105 1 : if (modelIt.second.find(communicator) != modelIt.second.end()) {
106 1 : modelIt.second.erase(communicator);
107 : }
108 : }
109 :
110 808 : HCCL_INFO("[%s] communicator[%p] resource release success", __func__, communicator);
111 808 : }
112 :
113 : // 记录aclgraph下发的所有tag, 首次记录时注册aclgraph销毁回调
114 7 : HcclResult AclgraphCallback::InsertNewTagToCaptureResMap(
115 : HcclCommunicator* communicator, const std::string& newTag, const OpParam& opParam)
116 : {
117 7 : CHK_PTR_NULL(communicator);
118 6 : aclmdlRI rtModel = nullptr;
119 6 : bool isCapture = false;
120 6 : u64 modelId = 0;
121 :
122 6 : CHK_RET(GetStreamCaptureInfo(opParam.stream.ptr(), rtModel, isCapture));
123 6 : CHK_PTR_NULL(rtModel);
124 6 : CHK_RET(GetModelId(rtModel, modelId));
125 :
126 6 : std::lock_guard<std::mutex> lock(resMutex_);
127 6 : if (captureResMap_.find(modelId) == captureResMap_.end()) {
128 4 : captureCallbackParamMap_[modelId].modelId = modelId;
129 8 : aclError aclRet = aclmdlRIDestroyRegisterCallback(
130 4 : rtModel, AclgraphDestroyCallback, static_cast<void*>(&captureCallbackParamMap_[modelId]));
131 4 : CHK_PRT_RET(
132 : aclRet != ACL_SUCCESS,
133 : HCCL_ERROR("[%s] aclmdlRIDestroyRegisterCallback fail, modelId[%llu]", __func__, modelId), HCCL_E_RUNTIME);
134 4 : HCCL_INFO("[%s] aclmdlRIDestroyRegisterCallback success modelID[%llu]", __func__, modelId);
135 : }
136 6 : captureResMap_[modelId][communicator].insert(newTag);
137 6 : HCCL_DEBUG("[%s] captureResMap insert tag[%s] to modelID[%llu]", __func__, newTag.c_str(), modelId);
138 :
139 6 : return HCCL_SUCCESS;
140 6 : }
141 : } // namespace hccl
|