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