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