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