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