Line data Source code
1 : /**
2 : * Copyright (c) 2025 Huawei Technologies Co., Ltd.
3 : * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4 : * CANN Open Software License Agreement Version 2.0 (the "License").
5 : * Please refer to the License for details. You may not use this file except in compliance with the License.
6 : * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7 : * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8 : * See LICENSE in the root of the software repository for the full text of the License.
9 : */
10 : #include "transport_urma_mem.h"
11 :
12 : namespace Hccl {
13 1 : TransportUrmaMem::TransportUrmaMem(BaseMemTransport *transport,
14 1 : RmaBufferMgr<BufferKey<uintptr_t, u64>, shared_ptr<HcclBuf>> &remoteHcclBufMgr)
15 1 : : transport_(transport), remoteHcclBufMgr_(remoteHcclBufMgr)
16 : {
17 1 : }
18 :
19 1 : TransportUrmaMem::~TransportUrmaMem()
20 : {
21 3 : HCCL_INFO("TransportUrmaMem Destroy");
22 1 : }
23 :
24 2 : HcclResult TransportUrmaMem::FillRmaBufferSlice(const RmaOpMem &localMem, const RmaOpMem &remoteMem,
25 : RmaBufferSlice& localRmaBufferSlice, RmtRmaBufferSlice& remoteRmaBufferSlice)
26 : {
27 2 : void* remoteAddr = remoteMem.addr;
28 2 : void* localAddr = localMem.addr;
29 2 : u64 byteSize = std::min(remoteMem.size, localMem.size);
30 2 : auto localKey = BufferKey<uintptr_t, u64>(reinterpret_cast<uintptr_t>(localAddr), byteSize);
31 2 : auto remoteKey = BufferKey<uintptr_t, u64>(reinterpret_cast<uintptr_t>(remoteAddr), byteSize);
32 :
33 2 : auto localBuffer = LocalUbRmaBufferManager::GetInstance()->Find(localKey);
34 2 : CHK_PRT_RET(!localBuffer.first,
35 : HCCL_ERROR("[TransportUrmaMem][FillRmaBufferSlice] Can't find localBuffer by key {%p, %llu}",
36 : localAddr, byteSize),
37 : HCCL_E_INTERNAL);
38 :
39 2 : auto remoteHcclBuf = remoteHcclBufMgr_.Find(remoteKey);
40 2 : CHK_PRT_RET(!remoteHcclBuf.first,
41 : HCCL_ERROR("[TransportUrmaMem][FillRmaBufferSlice] Can't find remoteBuffer by key {%p, %llu}",
42 : remoteAddr, byteSize),
43 : HCCL_E_INTERNAL);
44 2 : auto remoteBuffer = static_cast<RemoteUbRmaBuffer*>(remoteHcclBuf.second->handle);
45 :
46 2 : u64 localDataOffSet = static_cast<u8 *>(localAddr) - static_cast<u8 *>((void *)(localBuffer.second->GetBuf()->GetAddr()));
47 2 : u64 remoteDataOffSet = static_cast<u8 *>(remoteAddr) - static_cast<u8 *>(reinterpret_cast<void *>(remoteBuffer->GetAddr()));
48 :
49 2 : localRmaBufferSlice.addr = reinterpret_cast<u64>(static_cast<u8 *>((void *)(localBuffer.second->GetBuf()->GetAddr())) + localDataOffSet);
50 2 : localRmaBufferSlice.size = byteSize;
51 2 : localRmaBufferSlice.buf = localBuffer.second.get();
52 :
53 2 : remoteRmaBufferSlice.addr = reinterpret_cast<u64>(remoteBuffer->GetAddr() + remoteDataOffSet);
54 2 : remoteRmaBufferSlice.size = byteSize;
55 2 : remoteRmaBufferSlice.buf = remoteBuffer;
56 :
57 6 : HCCL_INFO("[TransportUrmaMem][FillRmaBufferSlice] Local [%p], buff[%lu], offset[%u], after mapping is [%llu], Datasize is [%llu].",
58 : localAddr, localBuffer.second->GetBuf()->GetAddr(), localDataOffSet, localRmaBufferSlice.addr, byteSize);
59 :
60 6 : HCCL_INFO("[TransportUrmaMem][FillRmaBufferSlice] rmt [%p], buff[%lu], offset[%u], after mapping is [%llu], Datasize is [%llu].",
61 : remoteAddr, remoteBuffer->GetAddr(), remoteDataOffSet, remoteRmaBufferSlice.addr, byteSize);
62 :
63 2 : return HCCL_SUCCESS;
64 2 : }
65 :
66 : // 2 is sizeof(float16), 8 is sizeof(float64), 2 is sizeof(bfloat16)..
67 : constexpr u32 SIZE_TABLE[HCCL_DATA_TYPE_RESERVED] = {sizeof(s8), sizeof(s16), sizeof(s32),
68 : 2, sizeof(float), sizeof(s64), sizeof(u64), sizeof(u8), sizeof(u16), sizeof(u32),
69 : 8, 2, 16, 2, 1, 1, 1, 1};
70 :
71 2 : inline HcclResult SalGetDataTypeSize(HcclDataType dataType, u32 &dataTypeSize)
72 : {
73 2 : if ((dataType >= HCCL_DATA_TYPE_INT8) &&
74 2 : (dataType < HCCL_DATA_TYPE_RESERVED)) {
75 2 : dataTypeSize = SIZE_TABLE[dataType];
76 : } else {
77 0 : HCCL_ERROR("[Get][DataTypeSize]errNo[0x%016llx] get date size failed. dataType[%u] is invalid.", \
78 : HCOM_ERROR_CODE(HcclResult::HCCL_E_PARA), dataType);
79 0 : return HCCL_E_PARA;
80 : }
81 2 : return HCCL_SUCCESS;
82 : }
83 :
84 2 : HcclResult TransportUrmaMem::BatchBufferSlice(const HcclOneSideOpDesc *oneSideDescs, u32 descNum,
85 : RmaBufferSlice *localRmaBufferSlice, RmtRmaBufferSlice *remoteRmaBufferSlice)
86 : {
87 6 : HCCL_INFO("[TransportUrmaMem][BatchBufferSlice] BatchBufferSlice Start");
88 :
89 : // 参数校验
90 2 : CHK_PTR_NULL(oneSideDescs);
91 2 : CHK_PTR_NULL(localRmaBufferSlice);
92 2 : CHK_PTR_NULL(remoteRmaBufferSlice);
93 :
94 2 : RmaOpMem remoteMem[MAX_DESC_NUM] = {};
95 2 : RmaOpMem localMem[MAX_DESC_NUM] = {};
96 :
97 2 : if (descNum > MAX_DESC_NUM) {
98 0 : THROW<InternalException>(StringFormat("[TransportUrmaMem][BatchBufferSlice] Desc item[%u] is out of range.", descNum));
99 : }
100 :
101 4 : for (u32 i = 0; i < descNum; i++) {
102 2 : if (oneSideDescs[i].count == 0) {
103 0 : HCCL_WARNING("[TransportUrmaMem][BatchBufferSlice] Desc item[%u] count is 0.", i);
104 : }
105 2 : u32 unitSize{0};
106 6 : HCCL_INFO("[TransportUrmaMem][BatchBufferSlice] SalGetDataTypeSize start");
107 2 : if (SalGetDataTypeSize(oneSideDescs[i].dataType, unitSize) != HCCL_SUCCESS) {
108 0 : THROW<InternalException>(StringFormat("[TransportUrmaMem][BatchBufferSlice] Get dataType size failed!"));
109 : }
110 :
111 2 : u64 byteSize = oneSideDescs[i].count * unitSize;
112 2 : remoteMem[i] = {oneSideDescs[i].remoteAddr, byteSize};
113 2 : localMem[i] = {oneSideDescs[i].localAddr, byteSize};
114 6 : HCCL_INFO("[TransportUrmaMem][BatchBufferSlice] FillRmaBufferSlice start");
115 2 : CHK_RET(FillRmaBufferSlice(localMem[i], remoteMem[i], localRmaBufferSlice[i], remoteRmaBufferSlice[i]));
116 : }
117 2 : return HCCL_SUCCESS;
118 : }
119 : } // namespace Hccl
|