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 "offload_stream_manager.h"
12 : #include "log.h"
13 : #include "exception_util.h"
14 : #include "invalid_params_exception.h"
15 : #include "stream_utils.h"
16 :
17 : namespace Hccl {
18 :
19 11 : void OffloadStreamManager::RegisterMaster(const std::string& opTag, std::unique_ptr<Stream> stream)
20 : {
21 33 : HCCL_INFO("[OffloadStreamManager::%s] start.", __func__);
22 :
23 11 : if (masters.find(opTag) != masters.end()) {
24 1 : std::string msg = StringFormat("master stream of op[%s] has been registered.", opTag.c_str());
25 1 : THROW<InvalidParamsException>(msg);
26 1 : }
27 : // 判断是否为acl graph零拷贝切图模式,判断标志为主流是否被捕获
28 10 : bool isCapture = false;
29 10 : rtModel_t rtModel = nullptr;
30 10 : auto ret = GetStreamCaptureInfo(stream->GetPtr(), rtModel, isCapture);
31 10 : if (ret != HCCL_SUCCESS) {
32 0 : THROW<InvalidParamsException>(
33 0 : StringFormat("[OffloadStreamManager::%s] GetStreamCaptureInfo failed.", __func__));
34 : }
35 10 : if (!isCapture) {
36 10 : ActivateSlaveStreams(opTag, stream.get()); // 不是acl graph则维持原流程
37 : }
38 10 : masters[opTag] = std::move(stream);
39 :
40 10 : currOpTag = opTag;
41 :
42 30 : HCCL_INFO("[OffloadStreamManager::%s] end.", __func__);
43 10 : }
44 :
45 10 : void OffloadStreamManager::ActivateSlaveStreams(const std::string& opTag, const Stream* masterStream)
46 : {
47 30 : HCCL_INFO("[OffloadStreamManager::%s] start.", __func__);
48 :
49 10 : const auto& slaveStreams = slaves[opTag];
50 10 : int slaveNum = slaveStreams.size();
51 10 : u32 mainStreamId = masterStream->GetId();
52 10 : auto& activeSlaveStreams = streamActiveManager_[mainStreamId];
53 11 : for (const auto& slave : slaveStreams) {
54 1 : u32 slaveId = slave->GetId();
55 1 : if (activeSlaveStreams.insert(slaveId).second) {
56 1 : HrtStreamActive(slave->GetPtr(), masterStream->GetPtr());
57 : }
58 : }
59 30 : HCCL_INFO("[OffloadStreamManager::%s] end, slaveNum[%d].", __func__, slaveNum);
60 10 : }
61 :
62 2 : void OffloadStreamManager::RegisterSlaves(const std::string& opTag, const std::vector<void*>& slaveStreams)
63 : {
64 6 : HCCL_INFO("[OffloadStreamManager::%s] start.", __func__);
65 :
66 2 : if (slaves.find(opTag) != slaves.end()) {
67 1 : std::string msg = StringFormat("slave streams of op[%s] has been registered.", opTag.c_str());
68 1 : THROW<InvalidParamsException>(msg);
69 1 : }
70 :
71 1 : int slaveNum = slaveStreams.size();
72 1 : slaves[opTag].resize(slaveNum);
73 3 : for (int i = 0; i < slaveNum; i++) {
74 2 : slaves[opTag][i] = std::make_unique<Stream>(slaveStreams[i], false);
75 : }
76 :
77 3 : HCCL_INFO("[OffloadStreamManager::%s] end, slaveNum[%d].", __func__, slaveNum);
78 1 : }
79 :
80 2 : Stream* OffloadStreamManager::GetSlave(const std::string& opTag)
81 : {
82 6 : HCCL_INFO("[OffloadStreamManager::%s] start, opTag[%s].", __func__, opTag.c_str());
83 :
84 2 : CheckOpTag(opTag);
85 :
86 2 : auto slavesIter = slaves.find(opTag);
87 2 : u32 slavesSize = slavesIter == slaves.end() ? 0 : slavesIter->second.size();
88 6 : HCCL_INFO("[OffloadStreamManager::%s] slavesSize[%u] slaveIndex[%u]", __func__, slavesSize, slaveIndex);
89 2 : if (slaveIndex >= slavesSize) {
90 0 : THROW<InvalidParamsException>(StringFormat("[OffloadStreamManager::%s] slave streams not enough.", __func__));
91 : }
92 :
93 6 : HCCL_INFO("[OffloadStreamManager::%s] end", __func__);
94 4 : return slaves[opTag][slaveIndex++].get();
95 : }
96 :
97 7 : Stream* OffloadStreamManager::GetMaster(const std::string& opTag)
98 : {
99 21 : HCCL_INFO("[OffloadStreamManager::%s] start, opTag[%s].", __func__, opTag.c_str());
100 :
101 7 : CheckOpTag(opTag);
102 :
103 7 : if (masters.find(opTag) == masters.end()) {
104 3 : HCCL_WARNING("[OffloadStreamManager::%s] master stream of opTag[%s] not found.", __func__, opTag.c_str());
105 1 : return nullptr;
106 : }
107 :
108 18 : HCCL_INFO("[OffloadStreamManager::%s] end", __func__);
109 6 : return masters[opTag].get();
110 : }
111 :
112 0 : u32 OffloadStreamManager::GetSlaveIndex(const std::string& opTag) const
113 : {
114 0 : CheckOpTag(opTag);
115 0 : return slaveIndex;
116 : }
117 :
118 1 : void OffloadStreamManager::ResetIndex(const std::string& opTag, u32 index)
119 : {
120 1 : CheckOpTag(opTag);
121 1 : slaveIndex = index;
122 1 : }
123 :
124 10 : void OffloadStreamManager::CheckOpTag(const std::string& opTag) const
125 : {
126 10 : if (opTag != currOpTag) {
127 0 : THROW<InvalidParamsException>(StringFormat(
128 : "[OffloadStreamManager::%s] opTag[%s] is not currOpTag[%s].", __func__, opTag.c_str(), currOpTag.c_str()));
129 : }
130 10 : }
131 :
132 0 : Stream* OffloadStreamManager::GetSlave(const std::string& opTag, u32 index) const
133 : {
134 0 : CheckOpTag(opTag);
135 0 : if (index >= slaves.at(opTag).size()) {
136 0 : THROW<InvalidParamsException>(
137 0 : StringFormat("[OffloadStreamManager::%s] index[%u] is invalid.", __func__, index));
138 : }
139 0 : return slaves.at(opTag)[index].get();
140 : }
141 :
142 1 : HcclResult OffloadStreamManager::ClearOpStream(const std::string& opTag)
143 : {
144 1 : if (masters.find(opTag) == masters.end()) {
145 3 : HCCL_WARNING("[OffloadStreamManager::%s] optag[%s] master stream not found.", __func__, opTag.c_str());
146 1 : return HCCL_SUCCESS;
147 : }
148 0 : if (slaves.find(opTag) == slaves.end()) {
149 0 : HCCL_WARNING("[OffloadStreamManager::%s] optag[%s] slave streams not found.", __func__, opTag.c_str());
150 0 : return HCCL_SUCCESS;
151 : }
152 0 : const auto& slaveStreams = slaves[opTag];
153 0 : u32 mainStreamId = masters[opTag]->GetId();
154 0 : auto& activeSlaveStreams = streamActiveManager_[mainStreamId];
155 0 : for (const auto& slave : slaveStreams) {
156 0 : u32 slaveId = slave->GetId();
157 0 : activeSlaveStreams.erase(slaveId);
158 : }
159 0 : masters.erase(opTag);
160 0 : slaves.erase(opTag);
161 0 : return HCCL_SUCCESS;
162 : }
163 :
164 : } // namespace Hccl
|