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

Generated by: LCOV version 2.0-1