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 : #include "comm_mems.h"
11 : #include <cstdlib>
12 : #include <algorithm>
13 : #include "orion_adapter_rts.h"
14 :
15 : namespace hccl {
16 :
17 183 : CommMemType ConvertHcclToCommMemType(HcclMemType hcclType)
18 : {
19 183 : switch (hcclType) {
20 36 : case HCCL_MEM_TYPE_DEVICE:
21 36 : return COMM_MEM_TYPE_DEVICE;
22 147 : case HCCL_MEM_TYPE_HOST:
23 147 : return COMM_MEM_TYPE_HOST;
24 0 : default:
25 0 : return COMM_MEM_TYPE_INVALID;
26 : }
27 : }
28 :
29 38 : HcclMemType ConvertCommToHcclMemType(CommMemType commType)
30 : {
31 38 : switch (commType) {
32 37 : case COMM_MEM_TYPE_DEVICE:
33 37 : return HCCL_MEM_TYPE_DEVICE;
34 1 : case COMM_MEM_TYPE_HOST:
35 1 : return HCCL_MEM_TYPE_HOST;
36 0 : default:
37 0 : return HCCL_MEM_TYPE_NUM;
38 : }
39 : }
40 :
41 178 : CommMems::CommMems(uint64_t bufferSize) : bufferSize_(bufferSize)
42 : {
43 178 : cclMemInfo_.mem.addr = nullptr;
44 178 : cclMemInfo_.mem.size = 0;
45 178 : cclMemInfo_.mem.type = CommMemType::COMM_MEM_TYPE_DEVICE;
46 178 : }
47 :
48 0 : HcclResult CommMems::Add(void* addr, uint64_t len) { return HCCL_SUCCESS; }
49 :
50 3 : HcclResult CommMems::GetHcclBuffer(void*& addr, uint64_t& len)
51 : {
52 3 : addr = reinterpret_cast<void*>(cclMemInfo_.mem.addr);
53 3 : len = static_cast<uint64_t>(cclMemInfo_.mem.size);
54 3 : return HCCL_SUCCESS;
55 : }
56 :
57 3 : HcclResult CommMems::HcclBufferMemset(void*& addr, uint64_t& len, bool clearFlag) const
58 : {
59 3 : if (!clearFlag) {
60 2 : HCCL_DEBUG("[CommMems][HcclBufferMemset] clearFlag[%d] is false, skip memset.", clearFlag);
61 2 : return HCCL_SUCCESS;
62 : }
63 :
64 1 : if (addr != nullptr && len > 0) {
65 1 : EXCEPTION_CATCH(Hccl::HrtMemset(addr, len, len), return HCCL_E_INTERNAL);
66 1 : return HCCL_SUCCESS;
67 : }
68 :
69 0 : HCCL_ERROR("[CommMems][HcclBufferMemset] buffer[%p] is null or size[%llu] is 0, skip.", addr, len);
70 0 : return HCCL_E_PARA;
71 : }
72 :
73 167 : HcclResult CommMems::Init(HcclMem cclBuffer)
74 : {
75 167 : cclMemInfo_.mem.addr = cclBuffer.addr;
76 167 : cclMemInfo_.mem.size = cclBuffer.size;
77 167 : cclMemInfo_.mem.type = ConvertHcclToCommMemType(cclBuffer.type);
78 167 : std::string memTag = "HcclBuffer";
79 167 : errno_t sRet = strncpy_s(cclMemInfo_.memTag, HCOMM_RES_TAG_MAX_LEN, memTag.c_str(), memTag.size());
80 167 : CHK_PRT_RET(sRet != EOK, HCCL_ERROR("[CommMems][Init] strncpy_s failed, return [%d].", sRet), HCCL_E_MEMORY);
81 167 : HCCL_INFO("[CommMems][Init] addr[%p] size[%llu] memType[%u]", cclBuffer.addr, cclBuffer.size, cclBuffer.type);
82 167 : memVersion_++;
83 167 : return HCCL_SUCCESS;
84 167 : }
85 :
86 48 : HcclResult CommMems::CommRegMem(const std::string& memTag, const CommMem& mem, void** memHandle)
87 : {
88 48 : CHK_PRT_RET(
89 : memHandle == nullptr, HCCL_ERROR("[CommRegMem] memHandle is null. tag[%s]", memTag.c_str()), HCCL_E_PARA);
90 47 : CHK_PRT_RET(
91 : mem.addr == nullptr || mem.size == 0,
92 : HCCL_ERROR("[CommRegMem] invalid mem. addr[%p] size[%llu]", mem.addr, (unsigned long long)mem.size),
93 : HCCL_E_PARA);
94 45 : if (UNLIKELY(memTag.size() >= HCOMM_RES_TAG_MAX_LEN)) {
95 1 : HCCL_ERROR("[CommRegMem] memTag.size() exceeds limit[%u]", HCOMM_RES_TAG_MAX_LEN);
96 1 : return HCCL_E_PARA;
97 : }
98 :
99 : // 组装句柄(仅域内管理,无进程级注册)
100 44 : Handle h;
101 44 : EXCEPTION_CATCH(h = std::make_shared<CommMemInfo>(), return HCCL_E_PTR);
102 44 : h->mem.addr = mem.addr;
103 44 : h->mem.size = mem.size;
104 44 : h->mem.type = mem.type;
105 44 : errno_t sRet = strncpy_s(h->memTag, HCOMM_RES_TAG_MAX_LEN, memTag.c_str(), memTag.size());
106 44 : CHK_PRT_RET(sRet != EOK, HCCL_ERROR("[CommRegMem] strncpy_s failed, return [%d].", sRet), HCCL_E_MEMORY);
107 :
108 44 : const auto key = MakeKey(mem.addr, static_cast<size_t>(mem.size));
109 :
110 44 : std::lock_guard<std::mutex> addLock(memMutex_);
111 :
112 44 : auto opIt = opBindings_.find(memTag);
113 44 : if (opIt != opBindings_.end()) {
114 2 : HCCL_ERROR(
115 : "[CommRegMem] memTag[%s] already registered: old addr[%p] size[%llu], new addr[%p] size[%llu].",
116 : memTag.c_str(), opIt->second->mem.addr, static_cast<unsigned long long>(opIt->second->mem.size), mem.addr,
117 : static_cast<unsigned long long>(mem.size));
118 2 : return HCCL_E_PARA;
119 : }
120 :
121 42 : auto& reg = tagRegs_[memTag];
122 :
123 : // 同tag内做区间冲突/幂等复用
124 42 : reg.table.AddWithoutCheck(key, h);
125 :
126 : // 加入绑定map
127 42 : opBindings_.emplace(memTag, h);
128 :
129 42 : memVersion_++;
130 42 : *memHandle = h.get();
131 42 : HCCL_INFO(
132 : "[CommRegMem] ok. tag[%s] memHandle[%p] size[%llu]", memTag.c_str(), *memHandle,
133 : static_cast<unsigned long long>(h->mem.size));
134 42 : return HCCL_SUCCESS;
135 44 : }
136 :
137 4 : HcclResult CommMems::CommUnregMem(const std::string& memTag, const void* memHandle) // 待确认是否要解注册
138 : {
139 4 : CHK_PRT_RET(memHandle == nullptr, HCCL_ERROR("[CommUnregMem] memHandle is null"), HCCL_E_PARA);
140 3 : CHK_PRT_RET(memTag.empty(), HCCL_ERROR("[CommUnregMem] memTag is null or empty"), HCCL_E_PARA);
141 :
142 2 : std::lock_guard<std::mutex> addLock(memMutex_);
143 :
144 2 : auto itTag = opBindings_.find(memTag);
145 2 : CHK_PRT_RET(
146 : itTag == opBindings_.end(), HCCL_WARNING("[CommUnregMem] tag[%s] not found in bindings", memTag.c_str()),
147 : HCCL_E_NOT_FOUND);
148 :
149 1 : auto& h = itTag->second; // Handle under this tag
150 1 : auto& reg = tagRegs_[itTag->first]; // TagRegistry for this tag
151 1 : size_t unboundCount = 0; // 本次解绑命中的句柄个数(即便 Del 未真正擦除也计数)
152 1 : size_t erasedCount = 0; // RmaBufferMgr::Del 返回 true 的次数(ref 归零而“擦除”)
153 :
154 1 : if (h.get() == memHandle) {
155 1 : const auto key = MakeKey(h->mem.addr, static_cast<size_t>(h->mem.size));
156 : try {
157 1 : if (reg.table.Del(key)) {
158 1 : ++erasedCount; // 该 key 的引用归零并从表中移除
159 : }
160 0 : } catch (const std::out_of_range&) {
161 0 : HCCL_ERROR("[CommUnregMem] tag[%s] key not found on Del (maybe already removed)", itTag->first.c_str());
162 0 : }
163 1 : ++unboundCount; // 从绑定列表移除,无论 Del 是否真正擦除
164 1 : opBindings_.erase(itTag);
165 1 : memVersion_++;
166 1 : if (reg.table.size() == 0) {
167 1 : tagRegs_.erase(std::string(memTag));
168 : }
169 : }
170 :
171 1 : CHK_PRT_RET(
172 : unboundCount == 0, HCCL_WARNING("[CommUnregMem] tag[%s] memHandle[%p] not found", memTag.c_str(), memHandle),
173 : HCCL_E_NOT_FOUND);
174 :
175 1 : HCCL_INFO(
176 : "[CommUnregMem] tag[%s] memHandle[%p] unbound=%zu, erased=%zu", memTag.c_str(), memHandle, unboundCount,
177 : erasedCount);
178 1 : return HCCL_SUCCESS;
179 2 : }
180 :
181 37 : HcclResult CommMems::GetTagsFromHandles(void** memHandles, uint32_t memHandleNum, std::vector<std::string>& memTags)
182 : {
183 37 : memTags.push_back("HcclBuffer");
184 :
185 37 : if (memHandles == nullptr) {
186 : // cclBuffer 已入 memTags,不因用户内存异常阻断注册流程,上游仍可获得 cclBuffer handle
187 32 : HCCL_INFO("[CommMems] memHandles is null, skip user memory");
188 32 : return HCCL_SUCCESS;
189 : }
190 :
191 5 : std::lock_guard<std::mutex> lock(memMutex_);
192 5 : CommMemInfo** handles = reinterpret_cast<CommMemInfo**>(memHandles);
193 11 : for (uint32_t i = 0; i < memHandleNum; i++) {
194 7 : if (handles[i] == nullptr) {
195 1 : HCCL_ERROR("[CommMems] memHandle[%p] not found", handles[i]);
196 1 : return HCCL_E_NOT_FOUND;
197 : }
198 12 : memTags.push_back(handles[i]->memTag);
199 : }
200 4 : return HCCL_SUCCESS;
201 5 : }
202 :
203 12 : HcclResult CommMems::GetAllMemory(std::vector<HcclMem>& memVec, std::vector<std::string>& memTags, uint64_t& version)
204 : {
205 : // cclBuffer
206 : HcclMem cclMem;
207 12 : cclMem.size = cclMemInfo_.mem.size;
208 12 : cclMem.type = ConvertCommToHcclMemType(cclMemInfo_.mem.type);
209 12 : cclMem.addr = cclMemInfo_.mem.addr;
210 12 : memVec.push_back(cclMem);
211 12 : memTags.push_back("HcclBuffer");
212 :
213 : // 遍历 opBindings_ 全量
214 12 : std::lock_guard<std::mutex> lock(memMutex_);
215 38 : for (auto& tagHandle : opBindings_) {
216 26 : auto& memInfo = tagHandle.second;
217 : HcclMem mem;
218 26 : mem.addr = memInfo->mem.addr;
219 26 : mem.size = memInfo->mem.size;
220 26 : mem.type = ConvertCommToHcclMemType(memInfo->mem.type);
221 26 : memVec.push_back(mem);
222 26 : memTags.push_back(tagHandle.first);
223 : }
224 :
225 12 : version = memVersion_;
226 12 : HCCL_INFO("[CommMems][%s] total memory count[%zu] version[%llu]", __func__, memVec.size(), version);
227 12 : return HCCL_SUCCESS;
228 12 : }
229 :
230 : } // namespace hccl
|