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 "task_abort_handler.h"
12 : #include "sal_pub.h"
13 : #include "hccl_communicator.h"
14 : #include "adapter_rts_common.h"
15 :
16 : using namespace hccl;
17 : using namespace std;
18 :
19 : static std::vector<HcclCommunicator *> commVector;
20 : static Referenced ref_;
21 : static std::mutex mutex_;
22 : struct TaskAbortCbArgs {
23 : u64 commVectorAddr;
24 : };
25 :
26 6 : TaskAbortHandler::TaskAbortHandler() {}
27 6 : TaskAbortHandler::~TaskAbortHandler() {}
28 :
29 8 : int32_t ProcessTaskAbortHandleCallback(int32_t deviceLogicId, aclrtDeviceTaskAbortStage stage, uint32_t timeout, void *args)
30 : {
31 8 : HcclUs startut = TIME_NOW();
32 8 : CHK_PTR_NULL(args);
33 8 : HCCL_INFO("ProcessTaskAbortHandleCallback begin, deviceLogicId [%d], stage [%d], args [%p], commVector v1 size [%u], ref_ count is [%d]",
34 : deviceLogicId, stage, args, commVector.size(), ref_.Count());
35 8 : const std::chrono::seconds localtimeout = std::chrono::seconds(timeout);
36 8 : HcclResult ret = HCCL_SUCCESS;
37 8 : if (localtimeout != std::chrono::seconds(0)) {
38 4 : if (stage == ACL_RT_DEVICE_TASK_ABORT_PRE) {
39 0 : for (size_t i = 0; i < commVector.size(); i++) {
40 0 : std::chrono::steady_clock::time_point startTime = std::chrono::steady_clock::now();
41 0 : ret = commVector[i]->Suspend();
42 0 : std::chrono::steady_clock::time_point curTime = std::chrono::steady_clock::now();
43 0 : if (ret != HCCL_SUCCESS && ret != HCCL_E_SUSPENDING) {
44 0 : HCCL_ERROR("[NsRecovery] finish suspend failed");
45 0 : return static_cast<int>(TaskAbortResult::TaskAbort_Fail);
46 : }
47 0 : HCCL_DEBUG("[NsRecovery]finish suspend success");
48 0 : const auto elapsed = std::chrono::duration_cast<std::chrono::seconds>(curTime - startTime);
49 0 : CHK_PRT_RET(elapsed > localtimeout,
50 : HCCL_ERROR("[NsRecovery][suspend] NsRecovery suspend timeOut"),
51 : static_cast<int>(TaskAbortResult::TaskAbort_TimeOut));
52 : }
53 4 : } else if (stage == ACL_RT_DEVICE_TASK_ABORT_POST) {
54 8 : for (size_t i = 0; i < commVector.size(); i++) {
55 4 : std::chrono::steady_clock::time_point startTime = std::chrono::steady_clock::now();
56 4 : ret = commVector[i]->StopExec();
57 4 : std::chrono::steady_clock::time_point curTime = std::chrono::steady_clock::now();
58 4 : if (ret != HCCL_SUCCESS && ret != HCCL_E_SUSPENDING) {
59 0 : HCCL_ERROR("[NsRecovery] finish stopExec failed");
60 0 : return static_cast<int>(TaskAbortResult::TaskAbort_Fail);
61 : }
62 4 : HCCL_DEBUG("[NsRecovery]finish stopExec success");
63 4 : const auto elapsed = std::chrono::duration_cast<std::chrono::seconds>(curTime - startTime);
64 4 : CHK_PRT_RET(elapsed > localtimeout,
65 : HCCL_ERROR("[NsRecovery][StopExec] NsRecovery StopExec timeOut"),
66 : static_cast<int>(TaskAbortResult::TaskAbort_TimeOut));
67 : }
68 4 : HcclUs stopExecUt = TIME_NOW();
69 4 : HCCL_RUN_INFO("TaskAbortHandler:ProcessTaskAbortHandleCallback, stopExec take time:[%lld]us",
70 : DURATION_US(stopExecUt - startut).count());
71 :
72 7 : for (size_t i = 0; i < commVector.size(); i++) {
73 4 : std::chrono::steady_clock::time_point startTime = std::chrono::steady_clock::now();
74 4 : ret = commVector[i]->Clean();
75 4 : std::chrono::steady_clock::time_point curTime = std::chrono::steady_clock::now();
76 4 : if (ret != HCCL_SUCCESS && ret != HCCL_E_SUSPENDING) {
77 1 : HCCL_ERROR("[NsRecovery] finish clean failed");
78 1 : return static_cast<int>(TaskAbortResult::TaskAbort_Fail);
79 : }
80 3 : HCCL_DEBUG("[NsRecovery]finish clean success");
81 3 : const auto elapsed = std::chrono::duration_cast<std::chrono::seconds>(curTime - startTime);
82 3 : CHK_PRT_RET(elapsed > localtimeout,
83 : HCCL_ERROR("[NsRecovery][clean] NsRecovery Clean timeOut"),
84 : static_cast<int>(TaskAbortResult::TaskAbort_TimeOut));
85 3 : CHK_RET(commVector[i]->Stop());
86 : }
87 : }
88 : } else {
89 4 : if (stage == ACL_RT_DEVICE_TASK_ABORT_PRE){
90 0 : for( size_t i = 0; i < commVector.size(); i++) {
91 0 : ret = commVector[i]->Suspend();
92 0 : if (ret != HCCL_SUCCESS && ret != HCCL_E_SUSPENDING) {
93 0 : HCCL_ERROR("[NsRecovery] finish suspend failed");
94 0 : return static_cast<int>(TaskAbortResult::TaskAbort_Fail);
95 : }
96 0 : HCCL_DEBUG("[NsRecovery]finish suspend success");
97 : }
98 4 : } else if (stage == ACL_RT_DEVICE_TASK_ABORT_POST) {
99 8 : for (size_t i = 0; i < commVector.size(); i++) {
100 4 : ret = commVector[i]->StopExec();
101 4 : if (ret != HCCL_SUCCESS && ret != HCCL_E_SUSPENDING) {
102 0 : HCCL_ERROR("[NsRecovery] finish stopExec failed");
103 1 : return static_cast<int>(TaskAbortResult::TaskAbort_Fail);
104 : }
105 4 : HCCL_DEBUG("[NsRecovery]finish stopExec success");
106 : }
107 4 : HcclUs stopExecUt = TIME_NOW();
108 4 : HCCL_RUN_INFO("TaskAbortHandler:ProcessTaskAbortHandleCallback, stopExec take time:[%lld]us",
109 : DURATION_US(stopExecUt - startut).count());
110 7 : for (size_t i = 0; i < commVector.size(); i++) {
111 4 : ret = commVector[i]->Clean();
112 4 : if (ret != HCCL_SUCCESS && ret != HCCL_E_SUSPENDING) {
113 1 : HCCL_ERROR("[NsRecovery] finish clean failed");
114 1 : return static_cast<int>(TaskAbortResult::TaskAbort_Fail);
115 : }
116 3 : HCCL_DEBUG("[NsRecovery]finish clean success");
117 3 : CHK_RET(commVector[i]->Stop());
118 : }
119 : }
120 : }
121 :
122 6 : HcclUs endut = TIME_NOW();
123 6 : HCCL_RUN_INFO(
124 : "TaskAbortHandler:ProcessTaskAbortHandleCallback, deviceLogicId [%d], stage [%d], total take time:[%lld]us",
125 : deviceLogicId, stage, DURATION_US(endut - startut).count());
126 6 : return static_cast<int>(TaskAbortResult::TaskAbort_Success);
127 : }
128 :
129 414 : HcclResult TaskAbortHandler::Init(HcclCommunicator * communicator)
130 : {
131 414 : std::unique_lock<std::mutex> lock(mutex_);
132 416 : HCCL_INFO("TaskAbortHandler::Init commVector size is [%d], ref_ count is [%d]", commVector.size(), ref_.Count());
133 416 : if (ref_.Count() == 0) {
134 250 : CHK_RET(hrtTaskAbortHandleCallback(ProcessTaskAbortHandleCallback, static_cast<void *>(&commVector)));
135 : }
136 416 : ref_.Ref();
137 416 : commVector.push_back(communicator);
138 :
139 416 : return HCCL_SUCCESS;
140 416 : }
141 :
142 608 : HcclResult TaskAbortHandler::DeInit(HcclCommunicator * communicator)
143 : {
144 608 : std::unique_lock<std::mutex> lock(mutex_);
145 609 : HCCL_INFO("TaskAbortHandler::DeInit commVector size is [%d], ref_ count is [%d]", commVector.size(), ref_.Count());
146 1384 : for (auto it = commVector.begin(); it != commVector.end();) {
147 775 : if (*it == communicator) {
148 412 : it = commVector.erase(it);
149 : } else {
150 363 : ++it;
151 : }
152 : }
153 609 : ref_.Unref();
154 609 : if (ref_.Count() == 0) {
155 249 : commVector.clear();
156 249 : CHK_RET(hrtTaskAbortHandleCallback(nullptr, nullptr));
157 : }
158 609 : return HCCL_SUCCESS;
159 609 : }
|