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 "symmetric_memory.h"
12 : #include <algorithm> // for std::max
13 : #include <cstddef>
14 : #include <exception>
15 : #include <list> // for SimpleVaAllocator
16 :
17 : namespace hccl
18 : {
19 : /**
20 : * @brief (内部) 简单的VA空间分配器
21 : */
22 : class SymmetricMemory::SimpleVaAllocator {
23 : std::list<FreeBlock> freeList_; // 按offset排序的空闲块
24 : std::mutex mutex_;
25 : size_t totalSize_;
26 :
27 : public:
28 82 : SimpleVaAllocator() : totalSize_(0) {}
29 82 : ~SimpleVaAllocator() {
30 82 : Destroy();
31 82 : }
32 :
33 : // 增加调试打印函数
34 4 : void Dump(const char* tag) {
35 4 : HCCL_ERROR("[%s] === VA Allocator Dump (Total: %zu) ===", tag, totalSize_);
36 4 : size_t freeSum = 0;
37 4 : int i = 0;
38 4 : constexpr double kPercentageMultiplier = 100.0;
39 10 : for (auto &block : freeList_) {
40 6 : HCCL_ERROR(" Block[%d]: offset %zu (0x%zx) -> size %zu (0x%zx) | end: %zu",
41 : i++, block.offset, block.offset, block.size, block.size, block.offset + block.size);
42 6 : freeSum += block.size;
43 : }
44 4 : HCCL_ERROR(" Total Free: %zu (%.2f%%)", freeSum, static_cast<double>(freeSum) / totalSize_ * kPercentageMultiplier);
45 4 : HCCL_ERROR("==========================================");
46 4 : }
47 :
48 40 : HcclResult Init(size_t size) {
49 40 : std::lock_guard<std::mutex> lock(mutex_);
50 40 : CHK_PRT_RET(size == 0, HCCL_ERROR("[SimpleVaAllocator][Init] invalid size: 0"), HCCL_E_PARA);
51 39 : totalSize_ = size;
52 39 : freeList_.push_back({0, size});
53 39 : return HCCL_SUCCESS;
54 40 : }
55 :
56 83 : void Destroy() {
57 83 : std::lock_guard<std::mutex> lock(mutex_);
58 83 : freeList_.clear();
59 83 : totalSize_ = 0;
60 83 : }
61 :
62 70 : HcclResult Reserve(size_t size, size_t align, size_t &offset) {
63 70 : std::lock_guard<std::mutex> lock(mutex_);
64 :
65 : // Debug: 打印请求信息
66 70 : HCCL_INFO("[VAAllocator] Request Reserve: size %zu, align %zu", size, align);
67 :
68 86 : for (auto it = freeList_.begin(); it != freeList_.end(); ++it) {
69 82 : size_t start = it->offset;
70 82 : size_t end = it->offset + it->size;
71 :
72 : // 计算对齐后的offset
73 82 : size_t alignedOffset = (start + align - 1) & ~(align - 1);
74 :
75 : // 检查对齐后的空间是否足够
76 82 : if (alignedOffset < end && (end - alignedOffset) >= size) {
77 : // 找到了
78 66 : offset = alignedOffset;
79 :
80 66 : size_t frontPad = alignedOffset - start;
81 66 : size_t backPad = (end) - (alignedOffset + size);
82 :
83 66 : HCCL_INFO("[VAAllocator] Found Block: [0x%zx, 0x%zx], Need aligned: 0x%zx. FrontPad: %zu, BackPad: %zu",
84 : start, end, alignedOffset, frontPad, backPad);
85 :
86 66 : auto to_erase = it;
87 : // 先插入后部碎片(如果存在)
88 66 : if (backPad > 0) {
89 : // 后部碎片应该插入在to_erase之后
90 62 : freeList_.insert(std::next(to_erase),
91 62 : {alignedOffset + size, backPad});
92 : }
93 :
94 : // 再插入前部碎片(如果存在)
95 66 : if (frontPad > 0) {
96 : // 前部碎片插入在to_erase之前
97 4 : freeList_.insert(to_erase, {start, frontPad});
98 : }
99 :
100 : // 最后删除原空闲块
101 66 : freeList_.erase(to_erase);
102 66 : return HCCL_SUCCESS;
103 : } else {
104 : // 只有当块看起来比较大但因为对齐无法满足时才打印,避免刷屏
105 16 : if (it->size >= size) {
106 1 : HCCL_DEBUG("[VAAllocator] Block [0x%zx, 0x%zx] size %zu skipped. AlignedOffset 0x%zx overlaps end or insufficient.",
107 : start, end, it->size, alignedOffset);
108 : }
109 : }
110 : }
111 :
112 : // 失败时打印当前内存布局,极大概率是碎片化导致
113 4 : HCCL_ERROR("[VAAllocator] Failed to reserve size %zu with align %zu. No suitable block found.", size, align);
114 4 : Dump("Reserve Failed");
115 :
116 4 : return HCCL_E_MEMORY;
117 70 : }
118 :
119 27 : HcclResult Release(size_t offset, size_t size) {
120 27 : std::lock_guard<std::mutex> lock(mutex_);
121 : // 边界检查
122 27 : if (offset + size > totalSize_) {
123 1 : HCCL_ERROR("[VAAllocator] Release out of range. off %zu + size %zu > total %zu", offset, size, totalSize_);
124 1 : return HCCL_E_PARA;
125 : }
126 :
127 : // 找到插入位置并合并
128 26 : auto it = freeList_.begin();
129 41 : while (it != freeList_.end() && it->offset < offset) {
130 15 : ++it;
131 : }
132 :
133 : // 检查重叠,直接报错
134 : // 与前一块重叠
135 26 : if (it != freeList_.begin()) {
136 8 : auto prevIt = std::prev(it);
137 8 : if (prevIt->offset <= offset && prevIt->offset + prevIt->size >= offset + size) { // 完全重叠表示释放空闲区域
138 0 : HCCL_WARNING("[VAAllocator] Releasing block[0x%zx, size %zu] is free", offset, size);
139 0 : return HCCL_SUCCESS;
140 : }
141 8 : if (prevIt->offset + prevIt->size > offset) {
142 0 : HCCL_ERROR("[VAAllocator] Releasing block[0x%zx, size %zu] overlaps with the previous block.", offset, size);
143 0 : return HCCL_E_PARA;
144 : }
145 : }
146 : // 与后一块重叠
147 26 : if (it != freeList_.end() && it->offset < offset + size) {
148 3 : if (it->offset <= offset && it->offset + it->size >= offset + size) { // 完全重叠表示释放空闲区域
149 1 : HCCL_WARNING("[VAAllocator] Releasing block[0x%zx, size %zu] is free", offset, size);
150 1 : return HCCL_SUCCESS;
151 : }
152 2 : HCCL_ERROR("[VAAllocator] Releasing block[0x%zx, size %zu] overlaps with the next block.", offset, size);
153 2 : return HCCL_E_PARA;
154 : }
155 :
156 : // 插入新释放的块
157 23 : auto newIt = freeList_.insert(it, {offset, size});
158 23 : HCCL_INFO("[VAAllocator] Releasing block[0x%zx, size %zu]", offset, size);
159 :
160 : // 尝试与后一块合并
161 46 : if (std::next(newIt) != freeList_.end()) {
162 18 : auto nextIt = std::next(newIt);
163 18 : if (newIt->offset + newIt->size == nextIt->offset) {
164 14 : newIt->size += nextIt->size;
165 14 : freeList_.erase(nextIt);
166 : }
167 : }
168 : // 尝试与前一块合并
169 23 : if (newIt != freeList_.begin()) {
170 8 : auto prevIt = std::prev(newIt);
171 8 : if (prevIt->offset + prevIt->size == newIt->offset) {
172 3 : prevIt->size += newIt->size;
173 3 : freeList_.erase(newIt);
174 : }
175 : }
176 23 : return HCCL_SUCCESS;
177 27 : }
178 : };
179 :
180 81 : SymmetricMemory::SymmetricMemory(u32 rank, u32 rankSize, size_t stride, std::shared_ptr<SymmetricMemoryAgent> symmetricMemoryAgent)
181 81 : : SymmetricMemory(rank, rankSize, stride, SymmetricMemoryMode::HCCS, std::move(symmetricMemoryAgent))
182 : {
183 81 : }
184 :
185 189 : SymmetricMemory::SymmetricMemory(u32 rank, u32 rankSize, size_t stride, SymmetricMemoryMode mode,
186 189 : std::shared_ptr<SymmetricMemoryAgent> symmetricMemoryAgent)
187 189 : : rank_(rank),
188 189 : rankSize_(rankSize),
189 189 : mode_(mode),
190 189 : stride_(stride),
191 189 : symmetricMemoryAgent_(std::move(symmetricMemoryAgent)),
192 378 : isSingleRank_(rankSize == 1)
193 : {
194 189 : if (mode_ != SymmetricMemoryMode::URMA) {
195 82 : vaAllocator_.reset(new (std::nothrow) SimpleVaAllocator());
196 : }
197 189 : remoteShareablePids.resize(rankSize_, 0);
198 189 : }
199 :
200 189 : SymmetricMemory::~SymmetricMemory()
201 : {
202 189 : HCCL_INFO("[SymmetricMemory][~SymmetricMemory] begin");
203 189 : std::vector<void*> winHandles;
204 189 : winHandles.reserve(windowMap_.size());
205 197 : for (const auto& pair : windowMap_) {
206 8 : winHandles.emplace_back(pair.first);
207 : }
208 197 : for (void* winHandle : winHandles) {
209 8 : if (mode_ == SymmetricMemoryMode::URMA) {
210 2 : HcclResult ret = DeregisterUrmaSymmetricMem(winHandle);
211 2 : if (ret != HCCL_SUCCESS) {
212 0 : HCCL_WARNING("[SymmetricMemory][~SymmetricMemory] deregister URMA window failed, "
213 : "winHandle[%p], ret[%d], continue cleanup.", winHandle, ret);
214 : }
215 : } else {
216 6 : (void)DeregisterSymmetricMem(winHandle);
217 : }
218 : }
219 189 : for (void* winHandle : singleRankUrmaWindows_) {
220 0 : HcclResult ret = hrtFree(winHandle);
221 0 : if (ret != HCCL_SUCCESS) {
222 0 : HCCL_WARNING("[SymmetricMemory][~SymmetricMemory] free single rank URMA window failed, "
223 : "winHandle[%p], ret[%d], continue cleanup.", winHandle, ret);
224 : }
225 : }
226 189 : singleRankUrmaWindows_.clear();
227 189 : windowMap_.clear();
228 189 : sortedWindows_.clear();
229 189 : memoryResourceMap_.clear();
230 189 : remoteMemMap_.clear();
231 189 : importAddrs_.clear();
232 :
233 189 : if (heapBase_) {
234 8 : if (aclrtReleaseMemAddress(heapBase_) != ACL_SUCCESS) {
235 0 : HCCL_ERROR("[SymmetricMemory][~SymmetricMemory] Failed to release symmetric heap VA: %p", heapBase_);
236 : }
237 : }
238 :
239 189 : HCCL_INFO("[SymmetricMemory][~SymmetricMemory] end");
240 189 : }
241 :
242 15 : HcclResult SymmetricMemory::EnsureInit() {
243 15 : std::call_once(init_flag_, [this]() {
244 8 : initResult_ = Init();
245 8 : });
246 15 : return initResult_;
247 : }
248 :
249 20 : HcclResult SymmetricMemory::Init()
250 : {
251 20 : CHK_SMART_PTR_NULL(vaAllocator_);
252 :
253 19 : isSingleRank_ = (rankSize_ == 1);
254 19 : CHK_PRT_RET(isSingleRank_, HCCL_INFO("[SymmetricMemory][Init] single rank communicator"), HCCL_SUCCESS);
255 18 : CHK_PRT_RET(stride_ == 0, HCCL_ERROR("[SymmetricMemory][Init] invalid stride: 0"), HCCL_E_PARA);
256 18 : size_t free = 0;
257 18 : size_t total = 0;
258 18 : aclError acl_ret = aclrtGetMemInfo(ACL_HBM_MEM_HUGE, &free, &total); // 获取当前进程总的物理内存大小
259 18 : CHK_PRT_RET(acl_ret != ACL_SUCCESS,
260 : HCCL_ERROR("[SymmetricMemory][Init] aclrtGetMemInfo failed, ret=[%d]", acl_ret), HCCL_E_INTERNAL);
261 18 : CHK_PRT_RET(stride_ > total,
262 : HCCL_ERROR("[SymmetricMemory][Init] Stride[%llu] is out of total[%llu].", stride_, total), HCCL_E_PARA);
263 :
264 18 : acl_ret = aclrtMemGetAllocationGranularity(&prop, ACL_RT_MEM_ALLOC_GRANULARITY_RECOMMENDED, &granularity_);
265 18 : CHK_PRT_RET(acl_ret != ACL_SUCCESS,
266 : HCCL_ERROR("[SymmetricMemory][Init] Get memory granularity failed, ret=[%d]", acl_ret), HCCL_E_INTERNAL);
267 :
268 17 : CHK_PRT_RET(granularity_ == 0, HCCL_ERROR("[SymmetricMemory][Init] Invalid memory granularity: 0"), HCCL_E_INTERNAL);
269 :
270 16 : CHK_PRT_RET(stride_ % granularity_ != 0,
271 : HCCL_ERROR("[SymmetricMemory][Init] Stride %llu is not a multiple of granularity %zu.", stride_, granularity_), HCCL_E_PARA);
272 :
273 15 : size_t totalHeapSize = static_cast<size_t>(stride_ * rankSize_);
274 15 : void* hintPtr = reinterpret_cast<void*>(targetStartTB);
275 :
276 : // 每个rank都预留一个总大小为 totalHeapSize 的VA空间。
277 15 : if (aclrtReserveMemAddressNoUCMemory(&heapBase_, totalHeapSize, 0, hintPtr, 0) != ACL_SUCCESS) {
278 0 : HCCL_ERROR("[SymmetricMemory][Init] aclrtReserveMemAddress failed to reserve %zu bytes. stride: %llu, rankSize: %u.",
279 : totalHeapSize, stride_, rankSize_);
280 0 : return HCCL_E_INTERNAL;
281 : }
282 : // 初始化VA分配器 (管理本地rank的stride_大小空间,即管理偏移量。
283 : // 这是一个集合调用,所有rank上的vaAllocator_状态将保持一致(前提是 SimpleVaAllocator 是确定性的)
284 15 : CHK_RET(vaAllocator_->Init(stride_));
285 :
286 14 : CHK_SMART_PTR_NULL(symmetricMemoryAgent_);
287 13 : CHK_RET(symmetricMemoryAgent_->Init());
288 12 : CHK_RET(GetAllRankPid());
289 :
290 9 : HCCL_INFO("[SymmetricMemory][Init] SymmetricMemory initialized. Rank[%u], Local Heap Base: %p, Stride: %llu, "
291 : "RankSize: %u, mode[%u].", rank_, heapBase_, stride_, rankSize_, static_cast<u32>(mode_));
292 :
293 9 : return HCCL_SUCCESS;
294 : }
295 :
296 12 : HcclResult SymmetricMemory::GetAllRankPid()
297 : {
298 12 : int32_t localPid{0}; // 当前进程号
299 12 : if (aclrtDeviceGetBareTgid(&localPid) != ACL_SUCCESS) {
300 1 : HCCL_ERROR("[SymmetricMemory][GetAllRankPid] Failed to get pid");
301 1 : return HCCL_E_DRV;
302 : }
303 11 : HCCL_INFO("[SymmetricMemory][GetAllRankPid] Local pid: %d.", localPid);
304 :
305 11 : CHK_RET(symmetricMemoryAgent_->ExchangeInfo(static_cast<void*>(&localPid), static_cast<void*>(remoteShareablePids.data()), sizeof(localPid)));
306 :
307 9 : std::string pidStr;
308 27 : for (u32 i = 0; i < remoteShareablePids.size(); i++) {
309 18 : pidStr += std::to_string(remoteShareablePids[i]);
310 18 : pidStr += "; ";
311 : }
312 9 : HCCL_INFO("[SymmetricMemory][GetAllRankPid] remote pids: %s", pidStr.c_str());
313 :
314 9 : return HCCL_SUCCESS;
315 9 : }
316 :
317 5 : void* SymmetricMemory::AllocSymmetricMem(size_t size)
318 : {
319 5 : void* devWin = nullptr;
320 5 : void *ptr = nullptr;
321 5 : HcommResult ret = HcommMemAlloc(&ptr, size);
322 5 : if (ret != HCCL_SUCCESS) {
323 1 : HCCL_ERROR("[SymmetricMemory][AllocSymmetricMem] HcommMemAlloc failed for size[%u].", size);
324 1 : return nullptr;
325 : }
326 :
327 4 : ret = RegisterSymmetricMem(ptr, size, &devWin);
328 4 : if (ret != HCCL_SUCCESS) {
329 1 : HCCL_ERROR("[SymmetricMemory][AllocSymmetricMem] RegisterSymmetricMem failed for ptr[%p], size[%u].", ptr, size);
330 1 : (void)HcommMemFree(ptr);
331 1 : return nullptr;
332 : }
333 3 : return devWin;
334 : }
335 :
336 1 : HcclResult SymmetricMemory::FreeSymmetricMem(void* devWin)
337 : {
338 1 : std::shared_ptr<SymmetricWindow> pWin = windowMap_[devWin];
339 1 : if (pWin == nullptr) {
340 1 : return HCCL_SUCCESS;
341 : }
342 :
343 0 : void* userPtr = pWin->userVa;
344 0 : CHK_RET(DeregisterSymmetricMem(devWin));
345 0 : CHK_RET(static_cast<HcclResult>(HcommMemFree(userPtr)));
346 0 : return HCCL_SUCCESS;
347 1 : }
348 :
349 10 : HcclResult SymmetricMemory::AddSymmetricWindow(std::shared_ptr<SymmetricWindow> &win)
350 : {
351 10 : CHK_SMART_PTR_NULL(win);
352 10 : CHK_RET(hrtMalloc(&win->devWin, sizeof(SymmetricWindow)));
353 9 : CHK_RET(hrtMemSyncCopy(win->devWin, sizeof(SymmetricWindow),
354 : win.get(), sizeof(SymmetricWindow), HcclRtMemcpyKind::HCCL_RT_MEMCPY_KIND_HOST_TO_DEVICE));
355 :
356 9 : sortedWindows_.push_back(win);
357 9 : std::sort(sortedWindows_.begin(), sortedWindows_.end(),
358 2 : [](const std::shared_ptr<SymmetricWindow>& a, const std::shared_ptr<SymmetricWindow>& b) {
359 4 : return (reinterpret_cast<uintptr_t>(a->userVa) < reinterpret_cast<uintptr_t>(b->userVa)) ||
360 2 : ((reinterpret_cast<uintptr_t>(a->userVa) == reinterpret_cast<uintptr_t>(b->userVa)) &&
361 4 : (a->userSize < b->userSize));
362 : });
363 :
364 9 : windowMap_[win->devWin] = win;
365 9 : return HCCL_SUCCESS;
366 : }
367 :
368 14 : std::shared_ptr<SymmetricWindow> SymmetricMemory::FindExactUrmaSymmetricWindow(void* userVa, size_t userSize) const
369 : {
370 14 : const uintptr_t newStart = reinterpret_cast<uintptr_t>(userVa);
371 14 : auto insertIt = std::lower_bound(sortedWindows_.begin(), sortedWindows_.end(), newStart,
372 6 : [](const std::shared_ptr<SymmetricWindow>& window, uintptr_t addr) {
373 6 : return reinterpret_cast<uintptr_t>(window->userVa) < addr;
374 : });
375 14 : for (; insertIt != sortedWindows_.end(); ++insertIt) {
376 2 : if (reinterpret_cast<uintptr_t>((*insertIt)->userVa) != newStart) {
377 0 : break;
378 : }
379 2 : if ((*insertIt)->userSize == userSize) {
380 2 : return *insertIt;
381 : }
382 : }
383 12 : return nullptr;
384 : }
385 :
386 12 : std::shared_ptr<SymmetricWindow> SymmetricMemory::FindContainingUrmaSymmetricWindow(void* userVa, size_t userSize,
387 : u64* offset) const
388 : {
389 12 : if (userVa == nullptr || offset == nullptr) {
390 0 : return nullptr;
391 : }
392 12 : CHK_PRT_RET(userSize == 0, HCCL_ERROR("[SymmetricMemory][FindContainingUrmaSymmetricWindow] Invalid size: 0."),
393 : nullptr);
394 :
395 12 : const uintptr_t requestStart = reinterpret_cast<uintptr_t>(userVa);
396 12 : const uintptr_t requestEnd = requestStart + userSize;
397 12 : CHK_PRT_RET(requestEnd < requestStart,
398 : HCCL_ERROR("[SymmetricMemory][FindContainingUrmaSymmetricWindow] window address overflow, userVa[%p], "
399 : "size[%zu].", userVa, userSize), nullptr);
400 :
401 12 : auto upper = std::upper_bound(sortedWindows_.begin(), sortedWindows_.end(), requestStart,
402 4 : [](uintptr_t addr, const std::shared_ptr<SymmetricWindow>& window) {
403 4 : return addr < reinterpret_cast<uintptr_t>(window->userVa);
404 : });
405 14 : for (auto it = upper; it != sortedWindows_.begin();) {
406 4 : --it;
407 4 : const uintptr_t winStart = reinterpret_cast<uintptr_t>((*it)->userVa);
408 4 : const uintptr_t winEnd = winStart + (*it)->userSize;
409 4 : if (requestStart >= winStart && requestEnd <= winEnd) {
410 2 : *offset = requestStart - winStart;
411 2 : return *it;
412 : }
413 : }
414 10 : return nullptr;
415 : }
416 :
417 10 : std::shared_ptr<SymmetricWindow> SymmetricMemory::FindOverlappingUrmaSymmetricWindow(void* userVa, size_t userSize) const
418 : {
419 10 : if (userVa == nullptr) {
420 0 : return nullptr;
421 : }
422 10 : CHK_PRT_RET(userSize == 0, HCCL_ERROR("[SymmetricMemory][FindOverlappingUrmaSymmetricWindow] Invalid size: 0."),
423 : nullptr);
424 :
425 10 : const uintptr_t requestStart = reinterpret_cast<uintptr_t>(userVa);
426 10 : const uintptr_t requestEnd = requestStart + userSize;
427 10 : CHK_PRT_RET(requestEnd < requestStart,
428 : HCCL_ERROR("[SymmetricMemory][FindOverlappingUrmaSymmetricWindow] window address overflow, userVa[%p], "
429 : "size[%zu].", userVa, userSize), nullptr);
430 :
431 10 : auto upper = std::upper_bound(sortedWindows_.begin(), sortedWindows_.end(), requestStart,
432 2 : [](uintptr_t addr, const std::shared_ptr<SymmetricWindow>& window) {
433 2 : return addr < reinterpret_cast<uintptr_t>(window->userVa);
434 : });
435 10 : for (auto it = upper; it != sortedWindows_.begin();) {
436 2 : --it;
437 2 : const uintptr_t prevStart = reinterpret_cast<uintptr_t>((*it)->userVa);
438 2 : const uintptr_t prevEnd = prevStart + (*it)->userSize;
439 2 : if (requestStart < prevEnd && requestEnd > prevStart) {
440 2 : return *it;
441 : }
442 : }
443 8 : if (upper != sortedWindows_.end()) {
444 0 : const uintptr_t nextStart = reinterpret_cast<uintptr_t>((*upper)->userVa);
445 0 : const uintptr_t nextEnd = nextStart + (*upper)->userSize;
446 0 : if (requestStart < nextEnd && requestEnd > nextStart) {
447 0 : return *upper;
448 : }
449 : }
450 8 : return nullptr;
451 : }
452 :
453 0 : HcclResult SymmetricMemory::CheckUrmaSymmetricWindowRange(void* userVa, size_t userSize) const
454 : {
455 0 : CHK_PTR_NULL(userVa);
456 0 : CHK_PRT_RET(userSize == 0,
457 : HCCL_ERROR("[SymmetricMemory][CheckUrmaSymmetricWindowRange] Invalid size: 0."), HCCL_E_PARA);
458 0 : const uintptr_t newStart = reinterpret_cast<uintptr_t>(userVa);
459 0 : const uintptr_t newEnd = newStart + userSize;
460 0 : CHK_PRT_RET(newEnd < newStart,
461 : HCCL_ERROR("[SymmetricMemory][CheckUrmaSymmetricWindowRange] window address overflow, userVa[%p], size[%zu].",
462 : userVa, userSize), HCCL_E_PARA);
463 :
464 : // sortedWindows_按userVa升序维护,只需检查插入点前后窗口即可判断重叠。
465 0 : auto insertIt = std::upper_bound(sortedWindows_.begin(), sortedWindows_.end(), newStart,
466 0 : [](uintptr_t addr, const std::shared_ptr<SymmetricWindow>& window) {
467 0 : return addr < reinterpret_cast<uintptr_t>(window->userVa);
468 : });
469 0 : if (insertIt != sortedWindows_.begin()) {
470 0 : auto prevIt = std::prev(insertIt);
471 0 : const uintptr_t prevStart = reinterpret_cast<uintptr_t>((*prevIt)->userVa);
472 0 : const uintptr_t prevEnd = prevStart + (*prevIt)->userSize;
473 0 : CHK_PRT_RET(newStart < prevEnd,
474 : HCCL_ERROR("[SymmetricMemory][CheckUrmaSymmetricWindowRange] window overlaps previous, userVa[%p], "
475 : "size[%zu], prevUserVa[%p], prevSize[%zu].", userVa, userSize, (*prevIt)->userVa,
476 : (*prevIt)->userSize), HCCL_E_PARA);
477 : }
478 0 : if (insertIt != sortedWindows_.end()) {
479 0 : const uintptr_t nextStart = reinterpret_cast<uintptr_t>((*insertIt)->userVa);
480 0 : CHK_PRT_RET(newEnd > nextStart,
481 : HCCL_ERROR("[SymmetricMemory][CheckUrmaSymmetricWindowRange] window overlaps next, userVa[%p], size[%zu], "
482 : "nextUserVa[%p], nextSize[%zu].", userVa, userSize, (*insertIt)->userVa,
483 : (*insertIt)->userSize), HCCL_E_PARA);
484 : }
485 0 : return HCCL_SUCCESS;
486 : }
487 :
488 8 : HcclResult SymmetricMemory::AddUrmaSymmetricWindow(std::shared_ptr<SymmetricWindow> &win)
489 : {
490 8 : CHK_SMART_PTR_NULL(win);
491 8 : const uintptr_t newStart = reinterpret_cast<uintptr_t>(win->userVa);
492 8 : auto insertIt = std::upper_bound(sortedWindows_.begin(), sortedWindows_.end(), newStart,
493 1 : [](uintptr_t addr, const std::shared_ptr<SymmetricWindow>& window) {
494 1 : return addr < reinterpret_cast<uintptr_t>(window->userVa);
495 : });
496 :
497 8 : HcclResult ret = hrtMalloc(&win->devWin, sizeof(SymmetricWindow));
498 8 : if (ret != HCCL_SUCCESS) {
499 0 : HCCL_ERROR("[SymmetricMemory][AddUrmaSymmetricWindow] alloc device window failed, ret[%d].", ret);
500 0 : return ret;
501 : }
502 8 : ret = hrtMemSyncCopy(win->devWin, sizeof(SymmetricWindow),
503 8 : win.get(), sizeof(SymmetricWindow), HcclRtMemcpyKind::HCCL_RT_MEMCPY_KIND_HOST_TO_DEVICE);
504 8 : if (ret != HCCL_SUCCESS) {
505 0 : HCCL_ERROR("[SymmetricMemory][AddUrmaSymmetricWindow] copy window to device failed, ret[%d].", ret);
506 0 : CHK_PRT(hrtFree(win->devWin));
507 0 : win->devWin = nullptr;
508 0 : return ret;
509 : }
510 :
511 8 : sortedWindows_.insert(insertIt, win);
512 8 : windowMap_[win->devWin] = win;
513 8 : return HCCL_SUCCESS;
514 : }
515 :
516 1 : HcclResult SymmetricMemory::DeleteSymmetricWindow(std::shared_ptr<SymmetricWindow> &win)
517 : {
518 1 : auto it = std::find_if(sortedWindows_.begin(), sortedWindows_.end(),
519 1 : [&win](const std::shared_ptr<SymmetricWindow>& w) {
520 1 : return w.get() == win.get();
521 : });
522 1 : if (it != sortedWindows_.end()) {
523 1 : CHK_PRT(hrtFree(win->devWin));
524 1 : windowMap_.erase(win->devWin);
525 1 : sortedWindows_.erase(it);
526 : }
527 :
528 1 : return HCCL_SUCCESS;
529 : }
530 :
531 9 : HcclResult SymmetricMemory::DeleteSymmetricWindow(void* devWin)
532 : {
533 9 : auto it = windowMap_.find(devWin);
534 9 : if (it != windowMap_.end()) {
535 8 : std::shared_ptr<SymmetricWindow> win = it->second;
536 8 : CHK_PRT(hrtFree(win->devWin));
537 8 : windowMap_.erase(it);
538 :
539 8 : auto vecIt = std::find_if(sortedWindows_.begin(), sortedWindows_.end(),
540 8 : [&win](const std::shared_ptr<SymmetricWindow>& w) {
541 8 : return w.get() == win.get();
542 : });
543 8 : if (vecIt != sortedWindows_.end()) {
544 8 : sortedWindows_.erase(vecIt);
545 : }
546 8 : }
547 :
548 9 : return HCCL_SUCCESS;
549 : }
550 :
551 15 : HcclResult SymmetricMemory::GetMemoryInfo(void* ptr, size_t size, void** baseUserVa, size_t* baseVaSize, aclrtDrvMemHandle* paHandle)
552 : {
553 15 : CHK_PTR_NULL(ptr);
554 14 : CHK_PRT_RET(size == 0, HCCL_ERROR("[SymmetricMemory][GetMemoryInfo] Invalid size: 0."), HCCL_E_PARA);
555 :
556 : // 打印当前注册请求的关键信息
557 13 : HCCL_INFO("[SymmetricMemory][GetMemoryInfo] Request: ptr=%p, size=%zu, granularity=%zu",
558 : ptr, size, granularity_);
559 :
560 13 : if(aclrtMemGetAddressRange(ptr, baseUserVa, baseVaSize) != 0) {
561 1 : HCCL_ERROR("[SymmetricMemory][GetMemoryInfo] aclrtMemGetAddressRange failed for ptr[%p], size[%zu]. ", ptr, size);
562 1 : return HCCL_E_PARA;
563 : }
564 12 : CHK_PTR_NULL(*baseUserVa);
565 11 : CHK_PRT_RET(*baseVaSize == 0, HCCL_ERROR("[SymmetricMemory][GetMemoryInfo] Invalid baseVaSize: 0."), HCCL_E_PARA);
566 10 : CHK_PRT_RET(*baseVaSize % granularity_ != 0,
567 : HCCL_ERROR("[SymmetricMemory][GetMemoryInfo] baseVaSize %u is not a multiple of granularity %zu.",
568 : *baseVaSize, granularity_), HCCL_E_PARA);
569 :
570 9 : if (aclrtMemRetainAllocationHandle(*baseUserVa, paHandle) != 0) {
571 0 : HCCL_ERROR("[SymmetricMemory][GetMemoryInfo] MemRetainAllocationHandle failed for ptr[%p], size[%zu]. ", ptr, size);
572 0 : return HCCL_E_PARA;
573 : }
574 9 : CHK_PTR_NULL(*paHandle);
575 :
576 9 : if (reinterpret_cast<uintptr_t>(ptr) + size > reinterpret_cast<uintptr_t>(*baseUserVa) + *baseVaSize) {
577 1 : HCCL_ERROR("[SymmetricMemory][GetMemoryInfo] ptr=%p size=%zu exceeds block [baseUserVa=%p, size=%zu]",
578 : ptr, size, *baseUserVa, *baseVaSize);
579 1 : return HCCL_E_PARA;
580 : }
581 :
582 8 : HCCL_INFO("[SymmetricMemory][GetMemoryInfo] Retained paHandle[%p] for baseUserVa[%p], baseVaSize[%zu]. Total Stride: %zu",
583 : *paHandle, *baseUserVa, *baseVaSize, stride_);
584 :
585 8 : return HCCL_SUCCESS;
586 : }
587 :
588 15 : HcclResult SymmetricMemory::RegisterSymmetricMem(void* ptr, size_t size, void** devWin)
589 : {
590 15 : CHK_RET(EnsureInit());
591 15 : if (isSingleRank_) {
592 0 : HCCL_INFO("[SymmetricMemory][RegisterSymmetricMem] single rank communicator");
593 0 : CHK_RET(hrtMalloc(devWin, sizeof(SymmetricWindow)));
594 0 : return HCCL_SUCCESS;
595 : }
596 15 : CHK_PTR_NULL(devWin);
597 15 : void* baseUserVa = nullptr;
598 15 : size_t baseVaSize = 0;
599 : aclrtDrvMemHandle paHandle;
600 15 : CHK_RET(GetMemoryInfo(ptr, size, &baseUserVa, &baseVaSize, &paHandle));
601 :
602 8 : std::shared_ptr<PaMappingInfo> paMapInfo;
603 8 : auto it = paMappingMap_.find(paHandle);
604 8 : if (it != paMappingMap_.end()) {
605 1 : paMapInfo = it->second;
606 1 : paMapInfo->refCount++;
607 1 : HCCL_INFO("PA handle[%p], refCount[%d]", paHandle, paMapInfo->refCount);
608 : }else {
609 7 : size_t offset = 0;
610 : // 使用 granularity_ (通常是2MB) 作为对齐参数
611 7 : if (vaAllocator_->Reserve( baseVaSize, granularity_, offset) != HCCL_SUCCESS) {
612 0 : HCCL_ERROR("[SymmetricMemory][RegisterSymmetricMem] Failed to reserve VA space. "
613 : "Req alignedSize: %zu (0x%zx), Align: %zu. Total Stride: %zu. "
614 : "Is fragmentation too high or stride too small?",
615 : baseVaSize, baseVaSize, granularity_, stride_);
616 0 : return HCCL_E_MEMORY;
617 : }
618 7 : EXCEPTION_CATCH((paMapInfo = std::make_shared<PaMappingInfo>()), return HCCL_E_PTR);
619 7 : paMapInfo->paHandle = paHandle;
620 7 : paMapInfo->origAllocBaseVa = baseUserVa;
621 7 : paMapInfo->origAllocSize = baseVaSize;
622 7 : paMapInfo->heapBaseOffset = offset;
623 7 : paMapInfo->refCount = 1;
624 7 : paMappingMap_.emplace(paHandle, paMapInfo);
625 : }
626 8 : std::shared_ptr<SymmetricWindow> pWin = nullptr;
627 8 : EXCEPTION_CATCH((pWin = std::make_shared<SymmetricWindow>()), return HCCL_E_PTR);
628 8 : pWin->userVa = baseUserVa;
629 8 : pWin->userSize = baseVaSize;
630 8 : pWin->baseVa = static_cast<uint8_t*>(heapBase_) + paMapInfo->heapBaseOffset;
631 8 : pWin->alignedHeapOffset = paMapInfo->heapBaseOffset;
632 8 : pWin->alignedSize = baseVaSize;
633 8 : pWin->localRank = rank_;
634 8 : pWin->rankSize = rankSize_;
635 8 : pWin->stride = stride_;
636 8 : pWin->paHandle = paHandle;
637 8 : pWin->mode = mode_;
638 8 : pWin->remoteMems = nullptr;
639 8 : pWin->remoteMemNum = 0;
640 8 : HcclResult ret = RegisterInternal(paHandle, paMapInfo->heapBaseOffset, baseVaSize);
641 8 : if (ret != HCCL_SUCCESS) {
642 2 : HCCL_ERROR("[SymmetricMemory] RegisterInternal Failed!");
643 2 : goto INTERNAL_ERROR;
644 : }
645 6 : ret = AddSymmetricWindow(pWin);
646 6 : if (ret != HCCL_SUCCESS) {
647 1 : HCCL_ERROR("[SymmetricMemory] AddSymmetricWindow Failed!");
648 1 : goto INTERNAL_ERROR;
649 : }
650 :
651 5 : *devWin = pWin->devWin;
652 5 : return HCCL_SUCCESS;
653 :
654 3 : INTERNAL_ERROR:
655 3 : if (paMapInfo->refCount == 1) {
656 3 : HCCL_ERROR("[SymmetricMemory] Releasing offset 0x%zx", paMapInfo->heapBaseOffset);
657 3 : (void)vaAllocator_->Release(paMapInfo->heapBaseOffset, baseVaSize);
658 3 : paMappingMap_.erase(paHandle);
659 : } else {
660 0 : paMapInfo->refCount--;
661 : }
662 3 : return ret;
663 8 : }
664 :
665 14 : HcclResult SymmetricMemory::TryReuseRegisteredUrmaWindow(void* ptr, size_t size, void** devWin, bool &reused) const
666 : {
667 14 : reused = false;
668 14 : std::shared_ptr<SymmetricWindow> exactWindow = FindExactUrmaSymmetricWindow(ptr, size);
669 14 : if (exactWindow != nullptr) {
670 2 : *devWin = exactWindow->devWin;
671 2 : reused = true;
672 2 : HCCL_INFO("[SymmetricMemory][RegisterUrmaSymmetricMem] symmetric window already registered, "
673 : "reuse devWin[%p], ptr[%p], size[%zu].", *devWin, ptr, size);
674 2 : return HCCL_SUCCESS;
675 : }
676 :
677 12 : u64 parentOffset = 0;
678 12 : std::shared_ptr<SymmetricWindow> parentWindow = FindContainingUrmaSymmetricWindow(ptr, size, &parentOffset);
679 12 : if (parentWindow != nullptr) {
680 2 : *devWin = parentWindow->devWin;
681 2 : reused = true;
682 2 : HCCL_INFO("[SymmetricMemory][RegisterUrmaSymmetricMem] symmetric window is subset of registered window, "
683 : "reuse parent devWin[%p], request ptr[%p], request size[%zu], parent ptr[%p], parent size[%zu], "
684 : "offset[%llu].", *devWin, ptr, size, parentWindow->userVa, parentWindow->userSize, parentOffset);
685 2 : return HCCL_SUCCESS;
686 : }
687 :
688 10 : std::shared_ptr<SymmetricWindow> overlapWindow = FindOverlappingUrmaSymmetricWindow(ptr, size);
689 10 : if (overlapWindow != nullptr) {
690 2 : HCCL_INFO("[SymmetricMemory][RegisterUrmaSymmetricMem] symmetric window overlaps registered window, "
691 : "register independently, request ptr[%p], request size[%zu], overlap ptr[%p], overlap size[%zu].",
692 : ptr, size, overlapWindow->userVa, overlapWindow->userSize);
693 : }
694 10 : return HCCL_SUCCESS;
695 14 : }
696 :
697 8 : HcclResult SymmetricMemory::InitUrmaRemoteMems(std::vector<CommMem> &remoteMems, CommMem **devRemoteMems) const
698 : {
699 8 : remoteMems.resize(rankSize_);
700 24 : for (CommMem &remoteMem : remoteMems) {
701 16 : remoteMem.type = COMM_MEM_TYPE_INVALID;
702 16 : remoteMem.addr = nullptr;
703 16 : remoteMem.size = 0;
704 : }
705 :
706 8 : HcclResult ret = hrtMalloc(reinterpret_cast<void **>(devRemoteMems), remoteMems.size() * sizeof(CommMem));
707 8 : if (ret != HCCL_SUCCESS) {
708 0 : HCCL_ERROR("[SymmetricMemory][RegisterUrmaSymmetricMem] alloc device remoteMems failed, size[%zu], ret[%d].",
709 : remoteMems.size() * sizeof(CommMem), ret);
710 0 : return ret;
711 : }
712 8 : ret = hrtMemSyncCopy(*devRemoteMems, remoteMems.size() * sizeof(CommMem),
713 8 : remoteMems.data(), remoteMems.size() * sizeof(CommMem),
714 : HcclRtMemcpyKind::HCCL_RT_MEMCPY_KIND_HOST_TO_DEVICE);
715 8 : if (ret != HCCL_SUCCESS) {
716 0 : HCCL_ERROR("[SymmetricMemory][RegisterUrmaSymmetricMem] copy remoteMems to device failed, ret[%d].", ret);
717 0 : CHK_PRT(hrtFree(*devRemoteMems));
718 0 : return ret;
719 : }
720 8 : return HCCL_SUCCESS;
721 : }
722 :
723 8 : void SymmetricMemory::FillUrmaSymmetricWindow(std::shared_ptr<SymmetricWindow> &win, void* ptr, size_t size,
724 : CommMem *devRemoteMems) const
725 : {
726 8 : win->userVa = ptr;
727 8 : win->userSize = size;
728 8 : win->baseVa = nullptr;
729 8 : win->alignedHeapOffset = 0;
730 8 : win->alignedSize = 0;
731 8 : win->localRank = 0;
732 8 : win->rankSize = rankSize_;
733 8 : win->stride = 0;
734 8 : win->paHandle = nullptr;
735 8 : win->mode = mode_;
736 8 : win->remoteMems = devRemoteMems;
737 8 : win->remoteMemNum = rankSize_;
738 8 : }
739 :
740 11 : HcclResult SymmetricMemory::RegisterUrmaSymmetricMem(void* ptr, size_t size, void** devWin)
741 : {
742 11 : CHK_PTR_NULL(devWin);
743 11 : CHK_PTR_NULL(ptr);
744 11 : CHK_PRT_RET(mode_ != SymmetricMemoryMode::URMA,
745 : HCCL_ERROR("[SymmetricMemory][RegisterUrmaSymmetricMem] invalid mode[%d].", mode_), HCCL_E_PARA);
746 11 : CHK_PRT_RET(size == 0, HCCL_ERROR("[SymmetricMemory][RegisterUrmaSymmetricMem] Invalid size: 0."), HCCL_E_PARA);
747 11 : CHK_PRT_RET(reinterpret_cast<uintptr_t>(ptr) + size < reinterpret_cast<uintptr_t>(ptr),
748 : HCCL_ERROR("[SymmetricMemory][RegisterUrmaSymmetricMem] address overflow, ptr[%p], size[%zu].", ptr, size),
749 : HCCL_E_PARA);
750 : // URMA模式先校验窗口范围,避免非法/重叠地址进入后续device资源申请。
751 11 : bool reused = false;
752 11 : CHK_RET(TryReuseRegisteredUrmaWindow(ptr, size, devWin, reused));
753 11 : if (reused) {
754 2 : return HCCL_SUCCESS;
755 : }
756 9 : if (isSingleRank_) {
757 1 : HCCL_INFO("[SymmetricMemory][RegisterUrmaSymmetricMem] single rank communicator");
758 1 : CHK_RET(hrtMalloc(devWin, sizeof(SymmetricWindow)));
759 1 : singleRankUrmaWindows_.emplace(*devWin);
760 1 : return HCCL_SUCCESS;
761 : }
762 :
763 8 : std::shared_ptr<SymmetricWindow> pWin = nullptr;
764 8 : EXCEPTION_CATCH((pWin = std::make_shared<SymmetricWindow>()), return HCCL_E_PTR);
765 :
766 : // remoteMems初始为invalid,待ChannelAcquire交换完成后按remoteRank回填。
767 8 : CommMem *devRemoteMems = nullptr;
768 8 : std::vector<CommMem> remoteMems;
769 8 : CHK_RET(InitUrmaRemoteMems(remoteMems, &devRemoteMems));
770 8 : FillUrmaSymmetricWindow(pWin, ptr, size, devRemoteMems);
771 :
772 8 : HcclResult ret = AddUrmaSymmetricWindow(pWin);
773 8 : if (ret != HCCL_SUCCESS) {
774 0 : HCCL_ERROR("[SymmetricMemory][RegisterUrmaSymmetricMem] AddUrmaSymmetricWindow failed, ret[%d].", ret);
775 0 : CHK_PRT(hrtFree(devRemoteMems));
776 0 : return ret;
777 : }
778 :
779 8 : *devWin = pWin->devWin;
780 8 : remoteMemMap_[pWin->devWin] = std::move(remoteMems);
781 8 : return HCCL_SUCCESS;
782 8 : }
783 :
784 5 : HcclResult SymmetricMemory::DeregisterSymmetricMem(void* devWin)
785 : {
786 5 : HcclResult ret = HCCL_SUCCESS;
787 5 : CHK_PTR_NULL(devWin);
788 3 : if (isSingleRank_) {
789 0 : HCCL_INFO("[SymmetricMemory][DeregisterSymmetricMem] single rank communicator");
790 0 : CHK_RET(hrtFree(devWin));
791 0 : return ret;
792 : }
793 :
794 3 : for (auto it = sortedWindows_.begin(); it != sortedWindows_.end();) {
795 2 : if ((*it)->devWin != devWin) {
796 0 : it++;
797 0 : continue;
798 : }
799 :
800 2 : std::shared_ptr<PaMappingInfo> paMapInfo = paMappingMap_[(*it)->paHandle];
801 2 : if (paMapInfo->refCount == 1) {
802 6 : for (u32 i = 0; i < rankSize_; i++) {
803 4 : void* virPtr = static_cast<uint8_t*>(heapBase_) + (stride_ * i) + (*it)->alignedHeapOffset;
804 4 : if (importAddrs_.find(virPtr) == importAddrs_.end()) {
805 0 : HCCL_ERROR("[SymmetricMemory][DeregisterSymmetricMem] Get paHandle failed for ptr[%p], rank[%u].", virPtr, i);
806 0 : ret = HCCL_E_INTERNAL;
807 0 : continue;
808 : }
809 4 : aclrtDrvMemHandle handle = importAddrs_[virPtr];
810 4 : HCCL_INFO("[SymmetricMemory][DeregisterSymmetricMem] Start to UnmapMem virPtr[%p], handle[%p], rank[%u].", virPtr, handle, i);
811 4 : aclError aclRet = aclrtUnmapMem(virPtr);
812 4 : if (aclRet != ACL_SUCCESS) {
813 0 : HCCL_ERROR("[SymmetricMemory][DeregisterSymmetricMem] Failed to unmap mem for rank %u at va %p, ret[%d].", i, virPtr, aclRet);
814 0 : ret = HCCL_E_DRV;
815 : }
816 4 : aclRet = aclrtFreePhysical(handle);
817 4 : if (aclRet != ACL_SUCCESS) {
818 2 : HCCL_ERROR("[SymmetricMemory][DeregisterSymmetricMem] Free Physical handle[%p] failed, ret[%d], rank[%u].", handle, aclRet, i);
819 2 : ret = HCCL_E_DRV;
820 : }
821 4 : importAddrs_.erase(virPtr);
822 : }
823 2 : vaAllocator_->Release((*it)->alignedHeapOffset, (*it)->alignedSize);
824 2 : paMappingMap_.erase((*it)->paHandle);
825 : } else {
826 0 : CHK_PRT_RET(aclrtFreePhysical((*it)->paHandle) != ACL_SUCCESS,
827 : HCCL_ERROR("[SymmetricMemory][DeregisterSymmetricMem] Free Physical handle[%p] failed.", (*it)->paHandle), HCCL_E_DRV);
828 0 : paMapInfo->refCount--;
829 : }
830 :
831 2 : it = sortedWindows_.erase(it);
832 2 : windowMap_.erase(devWin);
833 2 : CHK_RET(hrtFree(devWin));
834 2 : break;
835 2 : }
836 :
837 3 : return ret;
838 : }
839 :
840 11 : HcclResult SymmetricMemory::DeregisterUrmaSymmetricMem(void* devWin)
841 : {
842 11 : CHK_PTR_NULL(devWin);
843 11 : CHK_PRT_RET(mode_ != SymmetricMemoryMode::URMA,
844 : HCCL_ERROR("[SymmetricMemory][DeregisterUrmaSymmetricMem] invalid mode[%d].", mode_), HCCL_E_PARA);
845 11 : if (isSingleRank_) {
846 2 : HCCL_INFO("[SymmetricMemory][DeregisterUrmaSymmetricMem] single rank communicator");
847 2 : auto singleRankWinIt = singleRankUrmaWindows_.find(devWin);
848 2 : CHK_PRT_RET(singleRankWinIt == singleRankUrmaWindows_.end(),
849 : HCCL_ERROR("[SymmetricMemory][DeregisterUrmaSymmetricMem] Window handle[%p] is not registered.",
850 : devWin), HCCL_E_NOT_FOUND);
851 1 : CHK_RET(hrtFree(devWin));
852 1 : singleRankUrmaWindows_.erase(singleRankWinIt);
853 1 : return HCCL_SUCCESS;
854 : }
855 :
856 9 : auto winIt = windowMap_.find(devWin);
857 9 : CHK_PRT_RET(winIt == windowMap_.end(),
858 : HCCL_ERROR("[SymmetricMemory][DeregisterUrmaSymmetricMem] Window handle[%p] is not registered.", devWin),
859 : HCCL_E_NOT_FOUND);
860 :
861 8 : memoryResourceMap_.erase(devWin);
862 8 : remoteMemMap_.erase(devWin);
863 :
864 8 : std::shared_ptr<SymmetricWindow> win = winIt->second;
865 8 : if (win->remoteMems != nullptr) {
866 8 : CHK_PRT(hrtFree(win->remoteMems));
867 8 : win->remoteMems = nullptr;
868 8 : win->remoteMemNum = 0;
869 : }
870 8 : return DeleteSymmetricWindow(devWin);
871 8 : }
872 :
873 3 : HcclResult SymmetricMemory::FindSymmetricWindow(void* ptr, size_t size, void** win, u64 *offset)
874 : {
875 3 : CHK_PTR_NULL(ptr);
876 3 : CHK_PTR_NULL(win);
877 3 : CHK_PTR_NULL(offset);
878 3 : CHK_PRT_RET(isSingleRank_, HCCL_DEBUG("[SymmetricMemory][FindSymmetricWindow] single rank communicator"), HCCL_E_NOT_FOUND);
879 3 : uintptr_t userVaStart = reinterpret_cast<uintptr_t>(ptr);
880 3 : uintptr_t userVaEnd = userVaStart + size;
881 :
882 : // 遍历所有窗口
883 3 : for (const auto& pWin : sortedWindows_) {
884 2 : uintptr_t winStart = reinterpret_cast<uintptr_t>(pWin->userVa);
885 2 : if (winStart > userVaStart) {
886 2 : return HCCL_E_NOT_FOUND;
887 : }
888 :
889 1 : if (userVaStart >= winStart && userVaEnd <= winStart + pWin->userSize) {
890 1 : *win = pWin->devWin;
891 1 : *offset = userVaStart - winStart;
892 1 : return HCCL_SUCCESS;
893 : }
894 : }
895 :
896 1 : return HCCL_E_NOT_FOUND;
897 : }
898 :
899 1 : HcclResult SymmetricMemory::FindUrmaSymmetricWindow(void* ptr, size_t size, void** win, size_t *offset)
900 : {
901 1 : CHK_PTR_NULL(ptr);
902 1 : CHK_PTR_NULL(win);
903 1 : CHK_PTR_NULL(offset);
904 : // A5查询未命中不是错误,返回nullptr让算子侧按普通内存路径处理。
905 1 : *win = nullptr;
906 1 : *offset = 0;
907 1 : CHK_PRT_RET(mode_ != SymmetricMemoryMode::URMA,
908 : HCCL_ERROR("[SymmetricMemory][FindUrmaSymmetricWindow] invalid mode[%d].", mode_), HCCL_E_PARA);
909 1 : if (isSingleRank_) {
910 0 : HCCL_DEBUG("[SymmetricMemory][FindUrmaSymmetricWindow] single rank communicator");
911 0 : return HCCL_SUCCESS;
912 : }
913 :
914 1 : uintptr_t userVaStart = reinterpret_cast<uintptr_t>(ptr);
915 1 : uintptr_t userVaEnd = userVaStart + size;
916 1 : CHK_PRT_RET(userVaEnd < userVaStart,
917 : HCCL_ERROR("[SymmetricMemory][FindUrmaSymmetricWindow] address overflow, ptr[%p], size[%zu].", ptr, size),
918 : HCCL_E_PARA);
919 :
920 1 : auto upper = std::upper_bound(sortedWindows_.begin(), sortedWindows_.end(), userVaStart,
921 1 : [](uintptr_t addr, const std::shared_ptr<SymmetricWindow>& window) {
922 1 : return addr < reinterpret_cast<uintptr_t>(window->userVa);
923 : });
924 1 : if (upper == sortedWindows_.begin()) {
925 0 : return HCCL_SUCCESS;
926 : }
927 :
928 1 : const auto &candidate = *std::prev(upper);
929 1 : const uintptr_t winStart = reinterpret_cast<uintptr_t>(candidate->userVa);
930 1 : const uintptr_t winEnd = winStart + candidate->userSize;
931 1 : if (userVaStart >= winStart && userVaEnd <= winEnd) {
932 1 : *win = candidate->devWin;
933 1 : *offset = userVaStart - winStart;
934 : }
935 1 : return HCCL_SUCCESS;
936 : }
937 :
938 2 : HcclResult SymmetricMemory::GetPendingRegisterInfos(std::vector<SymmetricMemoryRegisterInfo> ®isterInfos) const
939 : {
940 2 : registerInfos.clear();
941 2 : if (mode_ != SymmetricMemoryMode::URMA || windowMap_.empty()) {
942 0 : return HCCL_SUCCESS;
943 : }
944 :
945 4 : for (const auto &winItem : windowMap_) {
946 2 : void *devWin = winItem.first;
947 2 : if (memoryResourceMap_.find(devWin) != memoryResourceMap_.end()) {
948 1 : continue;
949 : }
950 1 : const std::shared_ptr<SymmetricWindow> &win = winItem.second;
951 1 : CHK_SMART_PTR_NULL(win);
952 1 : registerInfos.push_back({devWin, win->userVa, win->userSize});
953 : }
954 2 : return HCCL_SUCCESS;
955 : }
956 :
957 3 : HcclResult SymmetricMemory::SetRegisteredMemoryResource(void* devWin, const SymmetricMemoryResource &resource)
958 : {
959 3 : CHK_PTR_NULL(devWin);
960 3 : CHK_PRT_RET(windowMap_.find(devWin) == windowMap_.end(),
961 : HCCL_ERROR("[SymmetricMemory][SetRegisteredMemoryResource] Window handle[%p] is not registered.", devWin),
962 : HCCL_E_NOT_FOUND);
963 3 : CHK_PRT_RET(resource.memHandle == nullptr || resource.memTag.empty(),
964 : HCCL_ERROR("[SymmetricMemory][SetRegisteredMemoryResource] invalid memory resource, devWin[%p], "
965 : "memHandle[%p], memTagEmpty[%d].", devWin, resource.memHandle, resource.memTag.empty()),
966 : HCCL_E_PARA);
967 :
968 3 : memoryResourceMap_[devWin] = resource;
969 3 : return HCCL_SUCCESS;
970 : }
971 :
972 1 : HcclResult SymmetricMemory::GetRegisteredMemoryResource(void* devWin, SymmetricMemoryResource &resource) const
973 : {
974 1 : CHK_PTR_NULL(devWin);
975 1 : auto resourceIt = memoryResourceMap_.find(devWin);
976 1 : CHK_PRT_RET(resourceIt == memoryResourceMap_.end(),
977 : HCCL_INFO("[SymmetricMemory][GetRegisteredMemoryResource] Window handle[%p] has no registered resource.",
978 : devWin), HCCL_E_NOT_FOUND);
979 :
980 1 : resource = resourceIt->second;
981 1 : return HCCL_SUCCESS;
982 : }
983 :
984 0 : void SymmetricMemory::RemoveRegisteredMemoryResource(void* devWin)
985 : {
986 0 : if (devWin == nullptr) {
987 0 : return;
988 : }
989 0 : memoryResourceMap_.erase(devWin);
990 : }
991 :
992 3 : static void AppendUniqueDirtyResource(std::vector<void*> &dirtyResources, void *devWin)
993 : {
994 3 : if (std::find(dirtyResources.begin(), dirtyResources.end(), devWin) == dirtyResources.end()) {
995 2 : dirtyResources.emplace_back(devWin);
996 : }
997 3 : }
998 :
999 2 : void SymmetricMemory::BuildRemoteMemTagIndex(std::unordered_map<std::string, std::vector<void*>> &tagIndex) const
1000 : {
1001 2 : tagIndex.clear();
1002 4 : for (const auto &resourceItem : memoryResourceMap_) {
1003 2 : const SymmetricMemoryResource &resource = resourceItem.second;
1004 2 : tagIndex[resource.memTag].emplace_back(resourceItem.first);
1005 : }
1006 2 : }
1007 :
1008 3 : HcclResult SymmetricMemory::UpdateRemoteMemForResource(uint32_t remoteRank, const CommMem &remoteMem, void *devWin,
1009 : const SymmetricMemoryResource &resource)
1010 : {
1011 3 : auto winIt = windowMap_.find(devWin);
1012 3 : auto remoteMemIt = remoteMemMap_.find(devWin);
1013 3 : CHK_PRT_RET(winIt == windowMap_.end() || remoteMemIt == remoteMemMap_.end(),
1014 : HCCL_ERROR("[SymmetricMemory][UpdateRemoteMem] window resource not found, win[%p], tag[%s].",
1015 : devWin, resource.memTag.c_str()), HCCL_E_NOT_FOUND);
1016 3 : std::shared_ptr<SymmetricWindow> win = winIt->second;
1017 3 : CHK_SMART_PTR_NULL(win);
1018 3 : CHK_PTR_NULL(win->remoteMems);
1019 3 : CHK_PRT_RET(remoteMem.addr == nullptr || remoteMem.size == 0,
1020 : HCCL_ERROR("[SymmetricMemory][UpdateRemoteMem] invalid remote symmetric mem, tag[%s], "
1021 : "remoteRank[%u], addr[%p], size[%llu].", resource.memTag.c_str(), remoteRank, remoteMem.addr,
1022 : remoteMem.size), HCCL_E_PARA);
1023 3 : CHK_PRT_RET(remoteMem.size != static_cast<uint64_t>(win->userSize),
1024 : HCCL_ERROR("[SymmetricMemory][UpdateRemoteMem] remote symmetric mem size mismatch, tag[%s], "
1025 : "remoteRank[%u], localSize[%zu], remoteSize[%llu].", resource.memTag.c_str(), remoteRank,
1026 : win->userSize, remoteMem.size), HCCL_E_PARA);
1027 3 : std::vector<CommMem> &windowRemoteMems = remoteMemIt->second;
1028 3 : CHK_PRT_RET(remoteRank >= windowRemoteMems.size(),
1029 : HCCL_ERROR("[SymmetricMemory][UpdateRemoteMem] remoteRank[%u] exceeds remoteMem size[%zu].",
1030 : remoteRank, windowRemoteMems.size()), HCCL_E_PARA);
1031 3 : windowRemoteMems[remoteRank] = remoteMem;
1032 3 : return HCCL_SUCCESS;
1033 3 : }
1034 :
1035 3 : HcclResult SymmetricMemory::UpdateRemoteMemByTag(uint32_t remoteRank, const CommMem &remoteMem, const char *memTag,
1036 : const std::unordered_map<std::string, std::vector<void*>> &tagIndex,
1037 : std::unordered_map<void*, bool> &matchedResources, std::vector<void*> &dirtyResources)
1038 : {
1039 3 : auto tagIt = tagIndex.find(memTag);
1040 3 : if (tagIt == tagIndex.end()) {
1041 0 : return HCCL_SUCCESS;
1042 : }
1043 6 : for (void *devWin : tagIt->second) {
1044 3 : auto resourceIt = memoryResourceMap_.find(devWin);
1045 3 : CHK_PRT_RET(resourceIt == memoryResourceMap_.end(),
1046 : HCCL_ERROR("[SymmetricMemory][UpdateRemoteMem] memory resource not found, win[%p], tag[%s].",
1047 : devWin, memTag), HCCL_E_NOT_FOUND);
1048 3 : const SymmetricMemoryResource &resource = resourceIt->second;
1049 3 : CHK_RET(UpdateRemoteMemForResource(remoteRank, remoteMem, devWin, resource));
1050 3 : matchedResources[devWin] = true;
1051 3 : AppendUniqueDirtyResource(dirtyResources, devWin);
1052 3 : HCCL_INFO("[SymmetricMemory][UpdateRemoteMem] record remote mem success, tag[%s], remoteRank[%u], "
1053 : "addr[%p], size[%llu].", resource.memTag.c_str(), remoteRank, remoteMem.addr, remoteMem.size);
1054 : }
1055 3 : return HCCL_SUCCESS;
1056 : }
1057 :
1058 2 : HcclResult SymmetricMemory::SyncDirtyRemoteMems(const std::vector<void*> &dirtyResources)
1059 : {
1060 4 : for (void *devWin : dirtyResources) {
1061 2 : auto winIt = windowMap_.find(devWin);
1062 2 : auto remoteMemIt = remoteMemMap_.find(devWin);
1063 2 : CHK_PRT_RET(winIt == windowMap_.end() || remoteMemIt == remoteMemMap_.end(),
1064 : HCCL_ERROR("[SymmetricMemory][UpdateRemoteMem] dirty window resource not found, win[%p].", devWin),
1065 : HCCL_E_NOT_FOUND);
1066 2 : std::shared_ptr<SymmetricWindow> win = winIt->second;
1067 2 : CHK_SMART_PTR_NULL(win);
1068 2 : CHK_PTR_NULL(win->remoteMems);
1069 2 : const std::vector<CommMem> &windowRemoteMems = remoteMemIt->second;
1070 2 : CHK_RET(hrtMemSyncCopy(win->remoteMems, windowRemoteMems.size() * sizeof(CommMem),
1071 : windowRemoteMems.data(), windowRemoteMems.size() * sizeof(CommMem),
1072 : HcclRtMemcpyKind::HCCL_RT_MEMCPY_KIND_HOST_TO_DEVICE));
1073 2 : HCCL_INFO("[SymmetricMemory][UpdateRemoteMem] sync remote mems success, win[%p], remoteMemNum[%zu].",
1074 : devWin, windowRemoteMems.size());
1075 2 : }
1076 2 : return HCCL_SUCCESS;
1077 : }
1078 :
1079 2 : HcclResult SymmetricMemory::CheckAllRemoteMemMatched(uint32_t remoteRank,
1080 : const std::unordered_map<void*, bool> &matchedResources) const
1081 : {
1082 4 : for (const auto &resourceItem : memoryResourceMap_) {
1083 2 : void *devWin = resourceItem.first;
1084 2 : const SymmetricMemoryResource &resource = resourceItem.second;
1085 2 : auto matchedIt = matchedResources.find(devWin);
1086 2 : CHK_PRT_RET(matchedIt == matchedResources.end() || !matchedIt->second,
1087 : HCCL_ERROR("[SymmetricMemory][UpdateRemoteMem] symmetric remote mem not found, tag[%s], "
1088 : "remoteRank[%u]. Please make sure all ranks register the same symmetric memory address and size.",
1089 : resource.memTag.c_str(), remoteRank), HCCL_E_PARA);
1090 : }
1091 2 : return HCCL_SUCCESS;
1092 : }
1093 :
1094 1 : HcclResult SymmetricMemory::UpdateRemoteMem(uint32_t remoteRank, const CommMem *remoteMems,
1095 : const std::vector<std::string> &memTags)
1096 : {
1097 1 : if (mode_ != SymmetricMemoryMode::URMA || memoryResourceMap_.empty()) {
1098 0 : return HCCL_SUCCESS;
1099 : }
1100 1 : CHK_PTR_NULL(remoteMems);
1101 1 : CHK_PRT_RET(remoteRank >= rankSize_,
1102 : HCCL_ERROR("[SymmetricMemory][UpdateRemoteMem] invalid remoteRank[%u], rankSize[%u].",
1103 : remoteRank, rankSize_), HCCL_E_PARA);
1104 :
1105 1 : std::unordered_map<void*, bool> matchedResources;
1106 2 : for (const auto &resourceItem : memoryResourceMap_) {
1107 1 : matchedResources.emplace(resourceItem.first, false);
1108 : }
1109 1 : std::unordered_map<std::string, std::vector<void*>> tagIndex;
1110 1 : BuildRemoteMemTagIndex(tagIndex);
1111 1 : std::vector<void*> dirtyResources;
1112 :
1113 2 : for (size_t memIdx = 0; memIdx < memTags.size(); ++memIdx) {
1114 1 : if (memTags[memIdx].empty()) {
1115 0 : continue;
1116 : }
1117 1 : CHK_RET(UpdateRemoteMemByTag(remoteRank, remoteMems[memIdx], memTags[memIdx].c_str(), tagIndex,
1118 : matchedResources, dirtyResources));
1119 : }
1120 1 : CHK_RET(SyncDirtyRemoteMems(dirtyResources));
1121 1 : return CheckAllRemoteMemMatched(remoteRank, matchedResources);
1122 1 : }
1123 :
1124 : // --- Private Methods ---
1125 8 : HcclResult SymmetricMemory::RegisterInternal(aclrtDrvMemHandle &paHandle, size_t offset, size_t mapSize)
1126 : {
1127 : aclrtMemFabricHandle shareableHandle;
1128 8 : if(paMappingMap_[paHandle]->refCount == 1) {
1129 7 : if (aclrtMemExportToShareableHandleV2(paHandle, 0,
1130 7 : ACL_MEM_SHARE_HANDLE_TYPE_FABRIC, static_cast<void*>(&shareableHandle)) != ACL_SUCCESS) {
1131 0 : HCCL_ERROR("[SymmetricMemory][RegisterInternal] Failed to export shareable handle. offset: %zu, size: %zu",
1132 : offset, mapSize);
1133 0 : return HCCL_E_DRV;
1134 : }
1135 :
1136 7 : if(aclrtMemSetPidToShareableHandleV2(static_cast<void*>(&shareableHandle), ACL_MEM_SHARE_HANDLE_TYPE_FABRIC,
1137 7 : remoteShareablePids.data(), remoteShareablePids.size()) != ACL_SUCCESS) {
1138 0 : HCCL_ERROR("[SymmetricMemory][RegisterInternal] Failed to aclrtMemSetPidToShareableHandleV2");
1139 0 : return HCCL_E_DRV;
1140 : }
1141 7 : paMappingMap_[paHandle]->shareableHandle = shareableHandle;
1142 : } else {
1143 1 : shareableHandle = paMappingMap_[paHandle]->shareableHandle;
1144 : }
1145 :
1146 8 : ShareableInfo shareableInfo{offset, mapSize, shareableHandle};
1147 8 : std::vector<ShareableInfo> remoteShareableInfos(rankSize_);
1148 :
1149 8 : CHK_RET(symmetricMemoryAgent_->ExchangeInfo(static_cast<void*>(&shareableInfo), static_cast<void*>(remoteShareableInfos.data()), sizeof(ShareableInfo)));
1150 23 : for (u32 i = 0; i < rankSize_; i++) {
1151 16 : if (remoteShareableInfos[i].offset != offset || remoteShareableInfos[i].size != mapSize) {
1152 1 : HCCL_ERROR("[SymmetricMemory][RegisterInternal] rank[%u]:[offset: %llu, mapSize: %llu] is not equal to "
1153 : "rank[%u]:[offset: %llu, mapSize: %llu]. Please ensure collective invocation!", rank_, offset, mapSize,
1154 : i, remoteShareableInfos[i].offset, remoteShareableInfos[i].size);
1155 1 : return HCCL_E_INTERNAL;
1156 : }
1157 : }
1158 :
1159 7 : u32 i = 0;
1160 7 : if(paMappingMap_[paHandle]->refCount == 1) {
1161 : aclrtDrvMemHandle importedHandle;
1162 16 : for (; i < rankSize_; i++) {
1163 11 : void* targetVa = static_cast<uint8_t*>(heapBase_) + (stride_ * i) + offset;
1164 11 : if (i == rank_) {
1165 6 : importedHandle = paHandle;
1166 5 : } else if (aclrtMemImportFromShareableHandleV2(static_cast<void*>(&remoteShareableInfos[i].handle), ACL_MEM_SHARE_HANDLE_TYPE_FABRIC, 0,
1167 5 : &importedHandle) != ACL_SUCCESS) {
1168 0 : HCCL_ERROR("[SymmetricMemory][RegisterInternal] Failed to import handle from rank %u.", i);
1169 1 : goto MAP_ERROR;
1170 : }
1171 :
1172 11 : if (aclrtMapMem(targetVa, mapSize, 0, importedHandle, 0) != ACL_SUCCESS) {
1173 1 : HCCL_ERROR("[SymmetricMemory][RegisterInternal] Failed to map mem for rank %u at va %p.", i, targetVa);
1174 1 : goto MAP_ERROR;
1175 : }
1176 10 : importAddrs_.insert({targetVa, importedHandle});
1177 10 : HCCL_INFO("[SymmetricMemory][RegisterInternal] success to Mapmem for rank %u at va %p to handle[%p].", i, targetVa, importedHandle);
1178 : }
1179 : }
1180 6 : return HCCL_SUCCESS;
1181 :
1182 1 : MAP_ERROR:
1183 1 : for (u32 j = 0; j < i; j++) {
1184 0 : (void)aclrtUnmapMem(static_cast<uint8_t*>(heapBase_) + (stride_ * j) + offset);
1185 : }
1186 1 : return HCCL_E_DRV;
1187 8 : }
1188 :
1189 : } // namespace hccl
|