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