Line data Source code
1 : /**
2 : * Copyright (c) 2026 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 <algorithm>
13 : #include <mutex>
14 : #include "log.h"
15 : #include "coll_comm.h"
16 : #include "ccu_device_pub.h"
17 :
18 : namespace hccl {
19 : using HcclUs = std::chrono::steady_clock::time_point;
20 :
21 5 : int32_t ProcessTaskAbortPre(const std::vector<CollComm*>& commVector, const std::chrono::seconds& localtimeout)
22 : {
23 5 : HcclResult ret = HCCL_SUCCESS;
24 5 : bool isUseTimeOut = localtimeout != std::chrono::seconds(0);
25 5 : std::chrono::seconds elapsed{};
26 8 : for (auto& comm : commVector) {
27 5 : if (isUseTimeOut) {
28 2 : std::chrono::steady_clock::time_point startTime = std::chrono::steady_clock::now();
29 2 : ret = comm->Suspend();
30 2 : elapsed = std::chrono::duration_cast<std::chrono::seconds>(std::chrono::steady_clock::now() - startTime);
31 : } else {
32 3 : ret = comm->Suspend();
33 : }
34 5 : if (ret != HCCL_SUCCESS && ret != HCCL_E_SUSPENDING) {
35 2 : HCCL_ERROR("[NsRecovery] finish suspend failed, ret = 0x%016llx", HCCL_ERROR_CODE(ret));
36 2 : return static_cast<int>(TaskAbortResult::TASK_ABORT_FAIL);
37 : }
38 3 : HCCL_INFO("[NsRecovery]finish suspend success");
39 3 : if (isUseTimeOut) {
40 1 : CHK_PRT_RET(
41 : elapsed > localtimeout, HCCL_ERROR("[NsRecovery][suspend] NsRecovery suspend timeOut"),
42 : static_cast<int>(TaskAbortResult::TASK_ABORT_TIMEOUT));
43 : }
44 : }
45 3 : return static_cast<int>(TaskAbortResult::TASK_ABORT_SUCCESS);
46 : }
47 :
48 8 : int32_t ProcessTaskAbortPost(
49 : const std::vector<CollComm*>& commVector, int32_t deviceLogicId, const std::chrono::seconds& localtimeout)
50 : {
51 8 : HcclResult ret = HCCL_SUCCESS;
52 8 : bool isUseTimeOut = localtimeout != std::chrono::seconds(0);
53 8 : std::chrono::seconds elapsed{};
54 8 : if (hcomm::CcuIsInited(deviceLogicId)) {
55 6 : CHK_RET(hcomm::CcuSetTaskKill(deviceLogicId));
56 : } else {
57 2 : HCCL_INFO("[NsRecovery][Callback] CCU not inited, skip CcuSetTaskKill, deviceLogicId[%d]", deviceLogicId);
58 : }
59 12 : for (auto& comm : commVector) {
60 7 : if (isUseTimeOut) {
61 3 : std::chrono::steady_clock::time_point startTime = std::chrono::steady_clock::now();
62 3 : ret = comm->Clean();
63 3 : elapsed = std::chrono::duration_cast<std::chrono::seconds>(std::chrono::steady_clock::now() - startTime);
64 : } else {
65 4 : ret = comm->Clean();
66 : }
67 7 : if (ret != HCCL_SUCCESS && ret != HCCL_E_SUSPENDING) {
68 2 : HCCL_ERROR("[NsRecovery][Callback] finish clean failed, ret = 0x%016llx", HCCL_ERROR_CODE(ret));
69 2 : return static_cast<int>(TaskAbortResult::TASK_ABORT_FAIL);
70 : }
71 5 : HCCL_INFO("[NsRecovery][Callback] finish clean success");
72 5 : if (isUseTimeOut) {
73 2 : CHK_PRT_RET(
74 : elapsed > localtimeout, HCCL_ERROR("[NsRecovery][Callback] NsRecovery Clean timeout"),
75 : static_cast<int>(TaskAbortResult::TASK_ABORT_TIMEOUT));
76 : }
77 : }
78 5 : if (hcomm::CcuIsInited(deviceLogicId)) {
79 3 : CHK_RET(hcomm::CcuSetTaskKillDone(deviceLogicId));
80 : } else {
81 2 : HCCL_INFO("[NsRecovery][Callback] CCU not inited, skip CcuSetTaskKillDone, deviceLogicId[%d]", deviceLogicId);
82 : }
83 5 : return static_cast<int>(TaskAbortResult::TASK_ABORT_SUCCESS);
84 : }
85 :
86 : int32_t
87 14 : ProcessTaskAbortHandleCallback(int32_t deviceLogicId, aclrtDeviceTaskAbortStage stage, uint32_t timeout, void* args)
88 : {
89 14 : HcclUs startut = std::chrono::steady_clock::now();
90 14 : CHK_PTR_NULL(args);
91 13 : auto* handler = static_cast<HcclTaskAbortHandler*>(args);
92 13 : std::lock_guard<std::mutex> lock(handler->vecMutex_);
93 13 : auto& commVector = handler->commVector_;
94 13 : HCCL_INFO("[NsRecovery][Callback] ProcessTaskAbortHandleCallback start!");
95 13 : const std::chrono::seconds localtimeout = std::chrono::seconds(timeout);
96 :
97 13 : if (stage == aclrtDeviceTaskAbortStage::ACL_RT_DEVICE_TASK_ABORT_PRE) {
98 5 : auto result = ProcessTaskAbortPre(commVector, localtimeout);
99 5 : if (result != static_cast<int>(TaskAbortResult::TASK_ABORT_SUCCESS)) {
100 2 : return result;
101 : }
102 8 : } else if (stage == aclrtDeviceTaskAbortStage::ACL_RT_DEVICE_TASK_ABORT_POST) {
103 8 : auto result = ProcessTaskAbortPost(commVector, deviceLogicId, localtimeout);
104 8 : if (result != static_cast<int>(TaskAbortResult::TASK_ABORT_SUCCESS)) {
105 3 : return result;
106 : }
107 : }
108 8 : HcclUs endut = std::chrono::steady_clock::now();
109 8 : auto execTime = std::chrono::duration_cast<std::chrono::microseconds>(endut - startut).count();
110 8 : HCCL_RUN_INFO("[NsRecovery][Callback] ProcessTaskAbortHandleCallback success, take time:[%lld]us", execTime);
111 8 : return static_cast<int>(TaskAbortResult::TASK_ABORT_SUCCESS);
112 13 : }
113 :
114 20 : HcclTaskAbortHandler::HcclTaskAbortHandler()
115 : {
116 20 : std::string name = "HCOMM";
117 20 : Hccl::HrtDeviceAbortRegCallBack(ProcessTaskAbortHandleCallback, static_cast<void*>(this), name);
118 20 : }
119 :
120 20 : HcclTaskAbortHandler::~HcclTaskAbortHandler()
121 : {
122 20 : std::string name = "HCOMM";
123 20 : Hccl::HrtDeviceAbortRegCallBack(nullptr, nullptr, name);
124 20 : }
125 :
126 186 : HcclResult HcclTaskAbortHandler::Register(CollComm* communicator)
127 : {
128 186 : std::lock_guard<std::mutex> lock(vecMutex_);
129 186 : commVector_.push_back(communicator);
130 186 : HCCL_INFO("HcclTaskAbortHandler::Register success, commVector_ size is [%zu]", commVector_.size());
131 :
132 186 : return HCCL_SUCCESS;
133 186 : }
134 :
135 192 : HcclResult HcclTaskAbortHandler::UnRegister(CollComm* communicator)
136 : {
137 192 : std::lock_guard<std::mutex> lock(vecMutex_);
138 192 : HCCL_INFO("HcclTaskAbortHandler::UnRegister Begin, commVector_ size is [%zu]", commVector_.size());
139 192 : auto it = std::find(commVector_.begin(), commVector_.end(), communicator);
140 192 : if (it != commVector_.end()) {
141 186 : commVector_.erase(it);
142 : } else {
143 6 : HCCL_WARNING("HcclTaskAbortHandler::UnRegister, comm not found.");
144 : }
145 192 : HCCL_INFO("HcclTaskAbortHandler::UnRegister finish, commVector_ size is [%zu]", commVector_.size());
146 192 : return HCCL_SUCCESS;
147 192 : }
148 : } // namespace hccl
|