LCOV - code coverage report
Current view: top level - legacy/ascend910/framework/communicator/impl/one_sided_service - hccl_one_sided_conn.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 292 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 19 0

            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 <numeric>
      12              : #include "sal_pub.h"
      13              : #include "hccl_one_sided_conn.h"
      14              : #include "p2p_mgmt_pub.h"
      15              : 
      16              : namespace hccl {
      17              : using namespace std;
      18              : 
      19            0 : HcclOneSidedConn::HcclOneSidedConn(
      20              :     const HcclNetDevCtx& netDevCtx, const HcclRankLinkInfo& localRankInfo, const HcclRankLinkInfo& remoteRankInfo,
      21              :     std::unique_ptr<HcclSocketManager>& socketManager, std::unique_ptr<NotifyPool>& notifyPool,
      22              :     const HcclDispatcher& dispatcher, const bool& useRdma, u32 sdid, u32 serverId, u32 trafficClass, u32 serviceLevel,
      23            0 :     bool aicpuUnfoldMode, bool isStandardCard, bool isNeedEnableP2P)
      24            0 :     : localRankInfo_(localRankInfo),
      25            0 :       socketManager_(socketManager),
      26            0 :       notifyPool_(notifyPool),
      27            0 :       aicpuUnfoldMode_(aicpuUnfoldMode),
      28            0 :       isStandardCard_(isStandardCard),
      29            0 :       isNeedEnableP2P_(isNeedEnableP2P)
      30              : {
      31            0 :     netDevCtx_ = netDevCtx;
      32            0 :     remoteRankInfo_ = remoteRankInfo;
      33            0 :     useRdma_ = useRdma;
      34            0 :     TransportMem::AttrInfo attrInfo{};
      35            0 :     attrInfo.localRankId = localRankInfo.userRank;
      36            0 :     attrInfo.remoteRankId = remoteRankInfo.userRank;
      37            0 :     attrInfo.sdid = sdid;
      38            0 :     attrInfo.serverId = serverId;
      39            0 :     attrInfo.trafficClass = trafficClass;
      40            0 :     attrInfo.serviceLevel = serviceLevel;
      41            0 :     if (useRdma) {
      42            0 :         transportMemPtr_ = TransportMem::Create(
      43            0 :             TransportMem::TpType::ROCE, notifyPool, netDevCtx, dispatcher, attrInfo, aicpuUnfoldMode_);
      44              :     } else {
      45            0 :         transportMemPtr_ = TransportMem::Create(
      46            0 :             TransportMem::TpType::IPC, notifyPool, netDevCtx, dispatcher, attrInfo, aicpuUnfoldMode_);
      47              :     }
      48            0 :     CHK_SMART_PTR_RET_NULL(transportMemPtr_);
      49            0 : }
      50              : 
      51            0 : HcclOneSidedConn::~HcclOneSidedConn()
      52              : {
      53            0 :     if ((isStandardCard_ && !useRdma_) && isNeedEnableP2P_) {
      54            0 :         if (!enableP2PDevices_.empty()) {
      55            0 :             P2PMgmtPub::DisableP2P(enableP2PDevices_);
      56            0 :             enableP2PDevices_.clear();
      57              :         }
      58              :     }
      59            0 : }
      60              : 
      61            0 : HcclResult HcclOneSidedConn::Connect(const std::string& commIdentifier, s32 timeoutSec)
      62              : {
      63            0 :     const auto startTime = TIME_NOW();
      64            0 :     if (aicpuUnfoldMode_) {
      65            0 :         CHK_RET(DeviceMem::alloc(transportDataDevice_, sizeof(TransportDeviceNormalData)));
      66              :     }
      67              :     // 创建socket用于交换数据
      68            0 :     std::string newTag;
      69            0 :     if (localRankInfo_.userRank < remoteRankInfo_.userRank) {
      70              :         // 本端为SERVER,对端为CLIENT
      71            0 :         newTag = string(localRankInfo_.ip.GetReadableIP()) + "_" + to_string(localRankInfo_.port) + "_"
      72            0 :                  + string(remoteRankInfo_.ip.GetReadableIP()) + "_" + to_string(remoteRankInfo_.port) + "_"
      73            0 :                  + commIdentifier;
      74              :     } else {
      75            0 :         newTag = string(remoteRankInfo_.ip.GetReadableIP()) + "_" + to_string(remoteRankInfo_.port) + "_"
      76            0 :                  + string(localRankInfo_.ip.GetReadableIP()) + "_" + to_string(localRankInfo_.port) + "_"
      77            0 :                  + commIdentifier;
      78              :     }
      79            0 :     HCCL_DEBUG("[HcclOneSidedConn][Connect]socket tag:%s", newTag.c_str());
      80              : 
      81              :     // 1、通信域初始化时会做非标卡且非310P场景的EnableP2P操作
      82              :     // 2、此处补全标卡且不使用RDMA场景下的EnableP2P操作
      83            0 :     HCCL_INFO(
      84              :         "[HcclOneSidedConn][Connect]localRankId[%u]-localDevicePhyId[%u], remoteRankId[%u]-remoteDevicePhyId[%u], "
      85              :         "isStandardCard[%s], useRdma[%s], isNeedEnableP2P[%s]",
      86              :         localRankInfo_.userRank, localRankInfo_.devicePhyId, remoteRankInfo_.userRank, remoteRankInfo_.devicePhyId,
      87              :         isStandardCard_ ? "true" : "false", useRdma_ ? "true" : "false", isNeedEnableP2P_ ? "true" : "false");
      88              : 
      89            0 :     if ((isStandardCard_ && !useRdma_) && isNeedEnableP2P_) {
      90            0 :         std::vector<u32> enableP2PDevices;
      91            0 :         enableP2PDevices.push_back(remoteRankInfo_.devicePhyId);
      92            0 :         HCCL_INFO(
      93              :             "[HcclOneSidedConn][Connect]localDevicePhyId[%u] enable p2p with remoteDevicePhyId[%u]",
      94              :             localRankInfo_.devicePhyId, remoteRankInfo_.devicePhyId);
      95            0 :         HcclResult ret = P2PMgmtPub::EnableP2P(enableP2PDevices);
      96            0 :         CHK_PRT_RET(
      97              :             ret != HCCL_SUCCESS,
      98              :             HCCL_ERROR(
      99              :                 "[HcclOneSidedConn][Connect]Enable P2P Failed, localPhyId[%u], remotephyId[%u], ret[%u]",
     100              :                 localRankInfo_.devicePhyId, remoteRankInfo_.devicePhyId, ret),
     101              :             ret);
     102            0 :         enableP2PDevices_.push_back(remoteRankInfo_.devicePhyId);
     103            0 :     }
     104              : 
     105              :     // EnableP2P需要和WaitP2PEnabled匹配使用,此处需要对1、2两处的EnableP2P做WaitP2PEnabled处理
     106            0 :     if ((!isStandardCard_ || !useRdma_) && isNeedEnableP2P_) {
     107            0 :         std::vector<u32> waitP2PEnabledDevices;
     108            0 :         waitP2PEnabledDevices.push_back(remoteRankInfo_.devicePhyId);
     109            0 :         HCCL_INFO(
     110              :             "[HcclOneSidedConn][Connect]localDevicePhyId[%u] wait p2p enable with remoteDevicePhyId[%u]",
     111              :             localRankInfo_.devicePhyId, remoteRankInfo_.devicePhyId);
     112            0 :         HcclResult ret = P2PMgmtPub::WaitP2PEnabled(waitP2PEnabledDevices, [this]() -> bool {
     113            0 :             return socketManager_->GetStopFlag();
     114              :         });
     115            0 :         CHK_PRT_RET(
     116              :             ret != HCCL_SUCCESS,
     117              :             HCCL_ERROR(
     118              :                 "[HcclOneSidedConn][Connect]Wait Enable P2P Failed, src devicePhyId[%u], dst devicePhyId[%u], ret[%u]",
     119              :                 localRankInfo_.devicePhyId, remoteRankInfo_.devicePhyId, ret),
     120              :             ret);
     121            0 :     }
     122              : 
     123            0 :     std::vector<std::shared_ptr<HcclSocket>> connectSockets;
     124            0 :     CHK_RET(socketManager_->CreateSingleLinkSocket(
     125              :         newTag, netDevCtx_, remoteRankInfo_, connectSockets, true, true, timeoutSec));
     126            0 :     CHK_RET(transportMemPtr_->SetDataSocket(connectSockets[0]));
     127            0 :     socket_ = connectSockets[0];
     128              : 
     129            0 :     if (useRdma_) {
     130              :         // 创建socket用于QP建链
     131            0 :         newTag += "_QP";
     132            0 :         auto timeCostSec = std::chrono::duration_cast<std::chrono::seconds>(TIME_NOW() - startTime).count();
     133            0 :         auto timeLeft = timeoutSec - timeCostSec;
     134            0 :         CHK_RET(socketManager_->CreateSingleLinkSocket(
     135              :             newTag, netDevCtx_, remoteRankInfo_, connectSockets, true, true, timeLeft));
     136            0 :         CHK_RET(transportMemPtr_->SetSocket(connectSockets[0]));
     137            0 :         rdmaSocket_ = connectSockets[0];
     138              : 
     139            0 :         if (timeoutSec == -1) {
     140              :             // timeout为-1,超时时间设为最大值
     141            0 :             CHK_RET(transportMemPtr_->Connect(INT_MAX));
     142              :         } else {
     143              :             // 超时时间减去已消耗的时间,避免接口整体耗时超过入参的秒数
     144            0 :             timeCostSec = std::chrono::duration_cast<std::chrono::seconds>(TIME_NOW() - startTime).count();
     145            0 :             timeLeft = timeoutSec - timeCostSec;
     146            0 :             CHK_PRT_RET(
     147              :                 timeLeft <= 0,
     148              :                 HCCL_ERROR(
     149              :                     "[HcclOneSidedConn][Connect] Connect timeout. comm[%s], timeoutSec[%d s]", commIdentifier.c_str(),
     150              :                     timeoutSec),
     151              :                 HCCL_E_TIMEOUT);
     152              : 
     153              :             // Transport建链:notify资源创建+QP建链
     154            0 :             CHK_RET(transportMemPtr_->Connect(timeLeft));
     155              :         }
     156              :     }
     157              : 
     158            0 :     return HCCL_SUCCESS;
     159            0 : }
     160              : 
     161            0 : void HcclOneSidedConn::CleanSocketResource(const std::string& commIdentifier)
     162              : {
     163              :     HcclSocketRole role;
     164            0 :     std::string newTag;
     165            0 :     if (localRankInfo_.userRank < remoteRankInfo_.userRank) {
     166              :         // 本端为SERVER,对端为CLIENT
     167            0 :         role = HcclSocketRole::SOCKET_ROLE_SERVER;
     168            0 :         newTag = string(localRankInfo_.ip.GetReadableIP()) + "_" + to_string(localRankInfo_.port) + "_"
     169            0 :                  + string(remoteRankInfo_.ip.GetReadableIP()) + "_" + to_string(remoteRankInfo_.port) + "_"
     170            0 :                  + commIdentifier;
     171              :     } else {
     172            0 :         role = HcclSocketRole::SOCKET_ROLE_CLIENT;
     173            0 :         newTag = string(remoteRankInfo_.ip.GetReadableIP()) + "_" + to_string(remoteRankInfo_.port) + "_"
     174            0 :                  + string(localRankInfo_.ip.GetReadableIP()) + "_" + to_string(localRankInfo_.port) + "_"
     175            0 :                  + commIdentifier;
     176              :     }
     177            0 :     if (socket_ != nullptr) {
     178            0 :         HCCL_INFO(
     179              :             "[HcclOneSidedConn][%s]abort and delete socket with remote[%u] tag[%s]", __func__, remoteRankInfo_.userRank,
     180              :             newTag.c_str());
     181            0 :         std::map<u32, std::vector<std::shared_ptr<HcclSocket>>> socketsMap;
     182            0 :         std::vector<std::shared_ptr<HcclSocket>> rankSockets{socket_};
     183            0 :         socketsMap.insert(std::make_pair(remoteRankInfo_.userRank, rankSockets));
     184            0 :         socketManager_->AbortAndDeleteSocket(newTag, role, socketsMap);
     185            0 :     }
     186              : 
     187            0 :     if (rdmaSocket_ != nullptr) {
     188            0 :         newTag += "_QP";
     189            0 :         HCCL_INFO(
     190              :             "[HcclOneSidedConn][%s]abort and delete rdmaSocket with remote[%u] tag[%s]", __func__,
     191              :             remoteRankInfo_.userRank, newTag.c_str());
     192            0 :         std::map<u32, std::vector<std::shared_ptr<HcclSocket>>> socketsMap;
     193            0 :         std::vector<std::shared_ptr<HcclSocket>> rankSockets{rdmaSocket_};
     194            0 :         socketsMap.insert(std::make_pair(remoteRankInfo_.userRank, rankSockets));
     195            0 :         socketManager_->AbortAndDeleteSocket(newTag, role, socketsMap);
     196            0 :     }
     197            0 :     return;
     198            0 : }
     199              : 
     200            0 : HcclResult HcclOneSidedConn::ExchangeIpcProcessInfo(const ProcessInfo& localProcess, ProcessInfo& remoteProcess)
     201              : {
     202            0 :     HCCL_DEBUG(
     203              :         "[HcclOneSidedConn][ExchangeIpcProcessInfo] localRank[%u] exchange process info", localRankInfo_.userRank);
     204            0 :     if (socket_->GetLocalRole() == HcclSocketRole::SOCKET_ROLE_CLIENT) {
     205            0 :         CHK_RET(socket_->Recv(&remoteProcess, sizeof(ProcessInfo)));
     206            0 :         CHK_RET(socket_->Send(&localProcess, sizeof(ProcessInfo)));
     207              :     } else {
     208            0 :         CHK_RET(socket_->Send(&localProcess, sizeof(ProcessInfo)));
     209            0 :         CHK_RET(socket_->Recv(&remoteProcess, sizeof(ProcessInfo)));
     210              :     }
     211            0 :     return HCCL_SUCCESS;
     212              : }
     213              : 
     214            0 : HcclResult HcclOneSidedConn::ExchangeMemDesc(
     215              :     const HcclMemDescs& localMemDescs, HcclMemDescs& remoteMemDescs, u32& actualNumOfRemote)
     216              : {
     217            0 :     TransportMem::RmaMemDesc* localMemDescArray
     218              :         = static_cast<TransportMem::RmaMemDesc*>(static_cast<void*>(localMemDescs.array));
     219            0 :     TransportMem::RmaMemDescs localRmaMemDescs = {localMemDescArray, localMemDescs.arrayLength};
     220            0 :     TransportMem::RmaMemDesc* remoteMemDescArray
     221              :         = static_cast<TransportMem::RmaMemDesc*>(static_cast<void*>(remoteMemDescs.array));
     222            0 :     TransportMem::RmaMemDescs remoteRmaMemDescs = {remoteMemDescArray, remoteMemDescs.arrayLength};
     223              : 
     224            0 :     return transportMemPtr_->ExchangeMemDesc(localRmaMemDescs, remoteRmaMemDescs, actualNumOfRemote);
     225              : }
     226              : 
     227            0 : HcclResult HcclOneSidedConn::GetMemType(const char* description, RmaMemType& memType)
     228              : {
     229            0 :     std::string tempDesc = std::string(description, TRANSPORT_EMD_ESC_SIZE);
     230            0 :     std::istringstream iss(tempDesc);
     231              :     // 定义需要跳过的变量的大小
     232              :     const std::vector<size_t> skip_sizes = {
     233              :         sizeof(u8),    // type
     234              :         sizeof(void*), // addr
     235              :         sizeof(u64),   // size
     236              :         sizeof(void*)  // devAddr
     237            0 :     };
     238              :     // 计算偏移量
     239            0 :     size_t offset = std::accumulate(skip_sizes.begin(), skip_sizes.end(), 0);
     240              :     // 定位到 memType 的位置
     241            0 :     iss.seekg(offset);
     242            0 :     iss.read(reinterpret_cast<char_t*>(&memType), sizeof(memType));
     243            0 :     CHK_PRT_RET(
     244              :         memType >= RmaMemType::TYPE_NUM,
     245              :         HCCL_ERROR("[HcclOneSidedConn][GetMemType] get memType failed memType[%d]", static_cast<int>(memType)),
     246              :         HCCL_E_INTERNAL);
     247            0 :     return HCCL_SUCCESS;
     248            0 : }
     249              : 
     250            0 : void HcclOneSidedConn::EnableMemAccess(const HcclMemDesc& remoteMemDesc, HcclMem& remoteMem)
     251              : {
     252              :     // 数据第一次转换
     253            0 :     HCCL_INFO("[HcclOneSidedConn][EnableMemAccess] Enable memory access.");
     254            0 :     const RmaMemDesc* remoteRmaMemDesc = static_cast<const RmaMemDesc*>(static_cast<const void*>(remoteMemDesc.desc));
     255            0 :     string tempStr = RmaMemDescCopyToStr(*remoteRmaMemDesc);
     256              :     RmaMemType memType;
     257            0 :     EXCEPTION_THROW_IF_ERR(
     258              :         GetMemType(remoteRmaMemDesc->memDesc, memType), "[HcclOneSidedConn][EnableMemAccess] get memType failed");
     259            0 :     auto iter = memDescMap_.find(tempStr);
     260            0 :     if (iter != memDescMap_.end()) {
     261            0 :         HcclBuf& outBuf = iter->second;
     262            0 :         BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(outBuf.addr), outBuf.len);
     263            0 :         auto resultPair = remoteRmaBufferMgr_.Add(tempKey, outBuf.handle);
     264            0 :         EXCEPTION_THROW_IF_COND_ERR(
     265              :             resultPair.first == remoteRmaBufferMgr_.End(),
     266              :             "[HcclOneSidedConn][EnableMemAccess]The memory that is expected to enable"
     267              :             " overlaps with the memory that has been enabled, please check params");
     268            0 :         remoteMem.type = static_cast<HcclMemType>(memType); // GE会检查这个字段先从字符串中获取
     269            0 :         remoteMem.addr = outBuf.addr;
     270            0 :         remoteMem.size = outBuf.len;
     271            0 :         return;
     272              :     }
     273              : 
     274              :     HcclBuf outBuf;
     275            0 :     EXCEPTION_THROW_IF_ERR(
     276              :         HcclMemImport(remoteRmaMemDesc->memDesc, TRANSPORT_EMD_ESC_SIZE, true, &outBuf, netDevCtx_),
     277              :         "[HcclOneSidedConn][EnableMemAccess] Enable memory access failed.");
     278            0 :     remoteMem.type = static_cast<HcclMemType>(memType); // GE会检查这个字段先从字符串中获取
     279            0 :     remoteMem.addr = outBuf.addr;
     280            0 :     remoteMem.size = outBuf.len;
     281            0 :     BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(outBuf.addr), outBuf.len);
     282            0 :     auto resultPair = remoteRmaBufferMgr_.Add(tempKey, outBuf.handle);
     283              : 
     284            0 :     EXCEPTION_THROW_IF_COND_ERR(
     285              :         resultPair.first == remoteRmaBufferMgr_.End(),
     286              :         "[HcclOneSidedConn][EnableMemAccess]The memory that is expected to enable"
     287              :         " overlaps with the memory that has been enabled, please check params");
     288            0 :     HCCL_INFO(
     289              :         "[HcclOneSidedConn][EnableMemAccess] after insert remoteRmaBufferMgr_ size[%d]", remoteRmaBufferMgr_.size());
     290              : 
     291            0 :     HCCL_INFO("[HcclOneSidedConn][EnableMemAccess] before insert memDescMap_ size[%d]", memDescMap_.size());
     292              : 
     293            0 :     memDescMap_.emplace(tempStr, outBuf);
     294            0 :     HCCL_INFO("[HcclOneSidedConn][EnableMemAccess] after insert memDescMap_ size[%d]", memDescMap_.size());
     295            0 :     HCCL_INFO("[HcclOneSidedConn][EnableMemAccess] Enable memory access success.");
     296            0 : }
     297              : 
     298            0 : void HcclOneSidedConn::DisableMemAccess(const HcclMemDesc& remoteMemDesc)
     299              : {
     300              :     // 数据第一次转换
     301            0 :     const RmaMemDesc* remoteRmaMemDesc = static_cast<const RmaMemDesc*>(static_cast<const void*>(remoteMemDesc.desc));
     302            0 :     string tempStr = RmaMemDescCopyToStr(*remoteRmaMemDesc);
     303            0 :     auto it = memDescMap_.find(tempStr);
     304            0 :     EXCEPTION_THROW_IF_COND_ERR(it == memDescMap_.end(), "Can't find hcclmem by key");
     305              : 
     306            0 :     BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(it->second.addr), it->second.len);
     307            0 :     HcclBuf& buf = it->second;
     308              :     try {
     309            0 :         if (remoteRmaBufferMgr_.Del(tempKey)) {
     310            0 :             EXCEPTION_THROW_IF_COND_ERR(HcclMemClose(&buf) != HCCL_SUCCESS, "Close remote memory failed.");
     311            0 :             HCCL_INFO("[HcclOneSidedConn][DisableMemAccess] before erase memDescMap_ size[%d]", memDescMap_.size());
     312            0 :             memDescMap_.erase(remoteRmaMemDesc->memDesc);
     313            0 :             HCCL_INFO("[HcclOneSidedConn][DisableMemAccess] after erase memDescMap_ size[%d]", memDescMap_.size());
     314              :             // 删除成功:输入key是表中某一最相近key的全集,计数-1后为0,返回true
     315            0 :             HCCL_INFO("[TransportIpcMem][DisableMemAccess]Memory reference count is 0, disable memory access.");
     316              :         } else {
     317              :             // 删除失败:输入key是表中某一最相近key的全集,计数不为0(存在其他remoteRank使用),返回false
     318            0 :             HCCL_INFO("[TransportIpcMem][DisableMemAccess]Memory reference count is larger than 0"
     319              :                       "(used by other RemoteRank), do not disable memory.");
     320              :         }
     321            0 :     } catch (std::out_of_range& e) {
     322            0 :         HCCL_ERROR("[TransportIpcMem][DisableMemAccess] catch RmaBufferMgr Del exception: %s", e.what());
     323            0 :         EXCEPTION_THROW_IF_COND_ERR(true, "[TransportIpcMem][DisableMemAccess] catch RmaBufferMgr Del exception");
     324            0 :     }
     325            0 :     HCCL_INFO("[HcclOneSidedConn][DisableMemAccess] Disable memory access success.");
     326            0 : }
     327              : 
     328            0 : void HcclOneSidedConn::BatchWrite(const HcclOneSideOpDesc* oneSideDescs, u32 descNum, const rtStream_t& stream)
     329              : {
     330            0 :     for (u32 i = 0; i < descNum; i++) {
     331            0 :         if (oneSideDescs[i].count == 0) {
     332            0 :             HCCL_WARNING("[HcclOneSidedConn][BatchWrite] Desc item[%u] count is 0.", i);
     333              :         }
     334              :         u32 unitSize;
     335            0 :         EXCEPTION_THROW_IF_ERR(
     336              :             SalGetDataTypeSize(oneSideDescs[i].dataType, unitSize),
     337              :             "[HcclOneSidedConn][BatchWrite] Get dataType size failed!");
     338            0 :         u64 byteSize = oneSideDescs[i].count * unitSize;
     339            0 :         HCCL_DEBUG(
     340              :             "[HcclOneSidedConn][BatchWrite] Desc[%u], localMem[%p], remoteMem[%p], size[%llu]", i,
     341              :             oneSideDescs[i].localAddr, oneSideDescs[i].remoteAddr, byteSize);
     342              : 
     343            0 :         BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(oneSideDescs[i].remoteAddr), byteSize);
     344              : 
     345            0 :         auto rmaBuffer = remoteRmaBufferMgr_.Find(tempKey);
     346            0 :         EXCEPTION_THROW_IF_COND_ERR(!rmaBuffer.first, "Can't find remoteBuffer by key");
     347              : 
     348            0 :         HcclBuf localMem = {oneSideDescs[i].localAddr, byteSize, nullptr};
     349            0 :         HcclBuf remoteMem = {oneSideDescs[i].remoteAddr, byteSize, rmaBuffer.second};
     350            0 :         EXCEPTION_THROW_IF_ERR(
     351              :             transportMemPtr_->Write(remoteMem, localMem, stream),
     352              :             "[HcclOneSidedConn][BatchWrite] transportMem WriteAsync failed.");
     353              :     }
     354            0 :     EXCEPTION_THROW_IF_ERR(transportMemPtr_->AddOpFence(stream), "[HcclOneSidedConn][BatchWrite] AddOpFence failed.");
     355            0 : }
     356              : 
     357            0 : void HcclOneSidedConn::BatchRead(const HcclOneSideOpDesc* oneSideDescs, u32 descNum, const rtStream_t& stream)
     358              : {
     359            0 :     for (u32 i = 0; i < descNum; i++) {
     360            0 :         if (oneSideDescs[i].count == 0) {
     361            0 :             HCCL_WARNING("[HcclOneSidedConn][BatchRead] Desc item[%u] count is 0.", i);
     362              :         }
     363              :         u32 unitSize;
     364            0 :         EXCEPTION_THROW_IF_ERR(
     365              :             SalGetDataTypeSize(oneSideDescs[i].dataType, unitSize),
     366              :             "[HcclOneSidedConn][BatchRead] Get dataType size failed!");
     367            0 :         u64 byteSize = oneSideDescs[i].count * unitSize;
     368            0 :         HCCL_DEBUG(
     369              :             "[HcclOneSidedConn][BatchRead] Desc[%u], localMem[%p], remoteMem[%p], size[%llu]", i,
     370              :             oneSideDescs[i].localAddr, oneSideDescs[i].remoteAddr, byteSize);
     371              : 
     372            0 :         BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(oneSideDescs[i].remoteAddr), byteSize);
     373              : 
     374            0 :         auto rmaBuffer = remoteRmaBufferMgr_.Find(tempKey);
     375            0 :         EXCEPTION_THROW_IF_COND_ERR(!rmaBuffer.first, "Can't find remoteBuffer by key");
     376              : 
     377            0 :         HcclBuf localMem = {oneSideDescs[i].localAddr, byteSize, nullptr};
     378            0 :         HcclBuf remoteMem = {oneSideDescs[i].remoteAddr, byteSize, rmaBuffer.second};
     379            0 :         EXCEPTION_THROW_IF_ERR(
     380              :             transportMemPtr_->Read(localMem, remoteMem, stream),
     381              :             "[HcclOneSidedConn][BatchRead] transportMem ReadAsync failed.");
     382              :     }
     383            0 :     EXCEPTION_THROW_IF_ERR(transportMemPtr_->AddOpFence(stream), "[HcclOneSidedConn][BatchRead] AddOpFence failed.");
     384            0 : }
     385              : 
     386            0 : HcclResult HcclOneSidedConn::GetTransInfo(
     387              :     HcclOneSideOpDescParam* descParam, const HcclOneSideOpDesc* desc, u32 descNum, u64& transportDataAddr,
     388              :     u64& transportDataSize)
     389              : {
     390            0 :     std::vector<u32> lkeys(descNum);
     391            0 :     std::vector<u32> rkeys(descNum);
     392            0 :     std::vector<HcclBuf> localMems(descNum);
     393            0 :     std::vector<HcclBuf> remoteMems(descNum);
     394            0 :     for (u32 i = 0; i < descNum - 1; ++i) { // last element is signal
     395            0 :         u32 unitSize = 0;
     396            0 :         CHK_RET(SalGetDataTypeSize(desc[i].dataType, unitSize));
     397            0 :         u64 bufSize = desc[i].count * unitSize;
     398            0 :         HCCL_DEBUG(
     399              :             "[HcclOneSidedConn][GetTransInfo] Desc[%u], localMem[%p], remoteMem[%p], size[%llu]", i, desc[i].localAddr,
     400              :             desc[i].remoteAddr, bufSize);
     401              : 
     402            0 :         BufferKey<uintptr_t, u64> bufKey(reinterpret_cast<uintptr_t>(desc[i].remoteAddr), bufSize);
     403            0 :         auto rmaBuffer = remoteRmaBufferMgr_.Find(bufKey);
     404            0 :         CHK_PRT_RET(
     405              :             !rmaBuffer.first,
     406              :             HCCL_ERROR("[GetTransInfo] Can't find remoteBuffer by key[%p][%llu]", desc[i].remoteAddr, bufSize),
     407              :             HCCL_E_PARA);
     408              : 
     409            0 :         localMems[i] = {desc[i].localAddr, bufSize, nullptr};
     410            0 :         remoteMems[i] = {desc[i].remoteAddr, bufSize, rmaBuffer.second};
     411              :     }
     412            0 :     CHK_RET(transportMemPtr_->GetTransInfo(
     413              :         transportData_.qpInfo, lkeys.data(), rkeys.data(), localMems.data(), remoteMems.data(), descNum));
     414            0 :     CHK_RET(hrtMemSyncCopy(
     415              :         transportDataDevice_.ptr(), transportDataDevice_.size(), reinterpret_cast<void*>(&transportData_),
     416              :         sizeof(transportData_), HcclRtMemcpyKind::HCCL_RT_MEMCPY_KIND_HOST_TO_DEVICE));
     417            0 :     transportDataAddr = reinterpret_cast<u64>(transportDataDevice_.ptr());
     418            0 :     transportDataSize = transportDataDevice_.size();
     419            0 :     for (u32 i = 0; i < descNum - 1; ++i) {
     420            0 :         u32 unitSize = 0;
     421            0 :         CHK_RET(SalGetDataTypeSize(desc[i].dataType, unitSize));
     422            0 :         descParam[i].dataType = static_cast<u8>(desc[i].dataType);
     423            0 :         descParam[i].count = localMems[i].len / unitSize;
     424            0 :         descParam[i].localAddr = reinterpret_cast<u64>(localMems[i].addr);
     425            0 :         descParam[i].remoteAddr = reinterpret_cast<u64>(remoteMems[i].addr);
     426            0 :         descParam[i].lkey = lkeys[i];
     427            0 :         descParam[i].rkey = rkeys[i];
     428              :     }
     429            0 :     descParam[descNum - 1].dataType = static_cast<u8>(HcclDataType::HCCL_DATA_TYPE_UINT8);
     430            0 :     descParam[descNum - 1].count = localMems[descNum - 1].len; // 因为dataType是UINT8,所以count等于len
     431            0 :     descParam[descNum - 1].localAddr = reinterpret_cast<u64>(localMems[descNum - 1].addr);
     432            0 :     descParam[descNum - 1].remoteAddr = reinterpret_cast<u64>(remoteMems[descNum - 1].addr);
     433            0 :     descParam[descNum - 1].lkey = lkeys[descNum - 1];
     434            0 :     descParam[descNum - 1].rkey = rkeys[descNum - 1];
     435            0 :     return HCCL_SUCCESS;
     436            0 : }
     437              : 
     438            0 : HcclResult HcclOneSidedConn::WaitOpFence(const rtStream_t& stream)
     439              : {
     440            0 :     CHK_RET(transportMemPtr_->WaitOpFence(stream));
     441            0 :     return HCCL_SUCCESS;
     442              : }
     443              : 
     444              : HcclResult
     445            0 : HcclOneSidedConn::ConnectWithRemote(const std::string& commIdentifier, ProcessInfo localProcess, s32 timeoutSec)
     446              : {
     447            0 :     CHK_RET(Connect(commIdentifier, timeoutSec));
     448            0 :     if (!useRdma_) {
     449            0 :         CHK_RET(ExchangeIpcProcessInfo(localProcess, remoteProcess_));
     450              :     }
     451            0 :     return HCCL_SUCCESS;
     452              : }
     453              : 
     454            0 : HcclResult HcclOneSidedConn::GetRemoteProcessInfo(ProcessInfo& remoteProcess)
     455              : {
     456            0 :     remoteProcess = remoteProcess_;
     457            0 :     return HCCL_SUCCESS;
     458              : }
     459              : 
     460            0 : HcclResult HcclOneSidedConn::ExchangeMemDesc(const HcclMemDescs& localMemDescs)
     461              : {
     462            0 :     constexpr u32 exchangeCntPerLoop = MAX_REMOTE_MEM_NUM;
     463            0 :     u32 localMemOffset = 0;
     464            0 :     u32 localMemCnt = localMemDescs.arrayLength;
     465            0 :     remoteMemDescsVec_.resize(exchangeCntPerLoop);
     466            0 :     actualNumOfRemote_ = 0;
     467              : 
     468              :     while (true) {
     469              :         // 每轮循环最多交换 exchangeCntPerLoop 个 memDesc
     470            0 :         u32 sendLocalCnt = localMemCnt > exchangeCntPerLoop ? exchangeCntPerLoop : localMemCnt;
     471            0 :         TransportMem::RmaMemDesc* localMemDescArray
     472            0 :             = static_cast<TransportMem::RmaMemDesc*>(static_cast<void*>(localMemDescs.array)) + localMemOffset;
     473            0 :         TransportMem::RmaMemDescs localRmaMemDescs = {localMemDescArray, sendLocalCnt};
     474              : 
     475            0 :         if (remoteMemDescsVec_.size() - actualNumOfRemote_ < exchangeCntPerLoop) {
     476            0 :             remoteMemDescsVec_.resize(remoteMemDescsVec_.size() + exchangeCntPerLoop);
     477              :         }
     478              :         TransportMem::RmaMemDesc* remoteMemDescArray
     479            0 :             = static_cast<TransportMem::RmaMemDesc*>(static_cast<void*>(&remoteMemDescsVec_[actualNumOfRemote_]));
     480            0 :         TransportMem::RmaMemDescs remoteRmaMemDescs = {remoteMemDescArray, exchangeCntPerLoop};
     481              : 
     482            0 :         u32 actualNumOfRemote = 0;
     483            0 :         CHK_RET(transportMemPtr_->ExchangeMemDesc(localRmaMemDescs, remoteRmaMemDescs, actualNumOfRemote));
     484            0 :         localMemOffset += sendLocalCnt;
     485            0 :         localMemCnt -= sendLocalCnt;
     486            0 :         actualNumOfRemote_ += actualNumOfRemote;
     487              : 
     488            0 :         if (actualNumOfRemote < exchangeCntPerLoop && sendLocalCnt < exchangeCntPerLoop) {
     489              :             // 循环结束条件,下轮没有memDesc要发 且 对端也没有memDesc要发
     490            0 :             break;
     491              :         }
     492            0 :     }
     493            0 :     return HCCL_SUCCESS;
     494              : }
     495              : 
     496            0 : HcclResult HcclOneSidedConn::EnableMemAccess()
     497              : {
     498            0 :     CHK_PRT_RET(
     499              :         actualNumOfRemote_ > remoteMemDescsVec_.size(),
     500              :         HCCL_ERROR(
     501              :             "[HcclOneSidedConn][EnableMemAccess] actualNumOfRemote[%u] is larger than remoteMemDescsVec.size[%zu]",
     502              :             actualNumOfRemote_, remoteMemDescsVec_.size()),
     503              :         HCCL_E_INTERNAL);
     504              : 
     505              :     HcclMem remoteMem;
     506            0 :     for (u32 i = 0; i < actualNumOfRemote_; i++) {
     507              :         // 创建HcclMemDesc对象
     508            0 :         HcclMemDesc* remoteMemDesc = static_cast<HcclMemDesc*>(static_cast<void*>(&remoteMemDescsVec_.at(i)));
     509            0 :         this->EnableMemAccess(*remoteMemDesc, remoteMem);
     510              :     }
     511            0 :     return HCCL_SUCCESS;
     512              : }
     513              : 
     514            0 : HcclResult HcclOneSidedConn::DisableMemAccess()
     515              : {
     516            0 :     CHK_PRT_RET(
     517              :         actualNumOfRemote_ > remoteMemDescsVec_.size(),
     518              :         HCCL_ERROR(
     519              :             "[HcclOneSidedConn][DisableMemAccess] actualNumOfRemote[%u] is larger than remoteMemDescsVec.size[%zu]",
     520              :             actualNumOfRemote_, remoteMemDescsVec_.size()),
     521              :         HCCL_E_INTERNAL);
     522              : 
     523            0 :     for (u32 i = 0; i < actualNumOfRemote_; i++) {
     524              :         // 创建HcclMemDesc对象
     525            0 :         HcclMemDesc* remoteMemDesc = static_cast<HcclMemDesc*>(static_cast<void*>(&remoteMemDescsVec_.at(i)));
     526            0 :         this->DisableMemAccess(*remoteMemDesc);
     527              :     }
     528              : 
     529            0 :     return HCCL_SUCCESS;
     530              : }
     531              : } // namespace hccl
        

Generated by: LCOV version 2.0-1