LCOV - code coverage report
Current view: top level - legacy/ascend910/framework/communicator/impl/symmetric_memory - symmetric_memory_agent.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 94.2 % 156 147
Test Date: 2026-08-18 17:47:01 Functions: 90.9 % 11 10

            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_agent.h"
      12              : #include <chrono>
      13              : 
      14              : namespace hccl {
      15              : using namespace std;
      16              : 
      17              : const string STR_IPC_MEM_EXCHANGE = "Exchange_Info";
      18              : constexpr u32 USLEEP_ONE_THOUSAND = 1000;
      19              : constexpr u32 RING_RANK_SIZE_MIN = 2;
      20              : 
      21           64 : SymmetricMemoryAgent::SymmetricMemoryAgent(
      22              :     const std::unique_ptr<HcclSocketManager>& socketManager, u32 devicePhyId, s32 deviceLogicId,
      23              :     const HcclIpAddress& localVnicIp, const std::vector<RankInfo>& rankInfoList, u32 userRank, bool useSuperPodMode,
      24           64 :     const std::string& identifier)
      25           64 :     : socketManager_(socketManager),
      26           64 :       devicePhyId_(devicePhyId),
      27           64 :       deviceLogicId_(deviceLogicId),
      28           64 :       localVnicIp_(localVnicIp),
      29           64 :       rankInfoList_(rankInfoList),
      30           64 :       userRank_(userRank),
      31          128 :       rankSize_(rankInfoList.size()),
      32           64 :       useSuperPodMode_(useSuperPodMode),
      33           64 :       identifier_(identifier)
      34              : {
      35           64 :     if (rankSize_ >= RING_RANK_SIZE_MIN) { // 当前数据交换算法使用超节点内大平面ring算法,需要和“左右”两边的rank建链
      36           61 :         leftRank_ = (userRank_ - 1 + rankSize_) % rankSize_;
      37           61 :         rightRank_ = (userRank_ + 1) % rankSize_;
      38              :     }
      39           64 : }
      40              : 
      41           64 : SymmetricMemoryAgent::~SymmetricMemoryAgent()
      42              : {
      43           64 :     threadRun_ = false;
      44           63 :     if (recvThread_ && recvThread_->joinable()) {
      45            9 :         recvThread_->join();
      46            9 :         recvThread_ = nullptr;
      47              :     }
      48           62 :     if (vnicPortCtx_ != nullptr) {
      49            8 :         HcclNetCloseDev(vnicPortCtx_);
      50            8 :         vnicPortCtx_ = nullptr;
      51              :     }
      52           62 : }
      53              : 
      54           10 : HcclResult SymmetricMemoryAgent::Init()
      55              : {
      56           10 :     CHK_PRT_RET(
      57              :         rankSize_ < RING_RANK_SIZE_MIN, HCCL_ERROR("[SymmetricMemoryAgent][Init] single rank communicator"),
      58              :         HCCL_E_PARA);
      59            9 :     CHK_RET(EstablishSockets());
      60            9 :     CHK_RET(InitRecvThread());
      61            9 :     return HCCL_SUCCESS;
      62              : }
      63              : 
      64            9 : HcclResult SymmetricMemoryAgent::InitRecvThread()
      65              : {
      66            9 :     threadRun_ = true;
      67            9 :     recvThread_.reset(new (std::nothrow) std::thread(&SymmetricMemoryAgent::DealWithRequest, std::ref(*this)));
      68            9 :     CHK_SMART_PTR_NULL(recvThread_);
      69            9 :     return HCCL_SUCCESS;
      70              : }
      71              : 
      72            8 : HcclResult SymmetricMemoryAgent::EstablishSockets()
      73              : {
      74            8 :     CHK_PRT_RET((vnicPortCtx_ != nullptr), HCCL_ERROR("[SymmetricMemoryAgent][Init] already initd"), HCCL_E_PARA);
      75            8 :     CHK_RET(HcclNetOpenDev(&vnicPortCtx_, NicType::VNIC_TYPE, devicePhyId_, deviceLogicId_, localVnicIp_));
      76            8 :     CHK_PTR_NULL(vnicPortCtx_);
      77              : 
      78            8 :     HCCL_INFO(
      79              :         "[SymmetricMemoryAgent][EstablishSockets] userRank[%u], leftRank_[%u], rightRank_[%u], rankSize_[%u]",
      80              :         userRank_, leftRank_, rightRank_, rankSize_);
      81           29 :     for (size_t i = 0; i < rankInfoList_.size(); i++) {
      82           21 :         if (rankInfoList_[i].userRank == leftRank_ || rankInfoList_[i].userRank == rightRank_) {
      83           10 :             HcclRankLinkInfo remoteLinkInfo;
      84           10 :             RankInfo dstRankInfo = rankInfoList_[i];
      85           10 :             remoteLinkInfo.userRank = dstRankInfo.userRank;
      86           10 :             remoteLinkInfo.devicePhyId = dstRankInfo.devicePhyId;
      87           10 :             remoteLinkInfo.ip = HcclIpAddress(dstRankInfo.devicePhyId);
      88           10 :             if (useSuperPodMode_) {
      89           10 :                 CHK_RET(hrtRaGetSingleSocketVnicIpInfo(
      90              :                     devicePhyId_, DeviceIdType::DEVICE_ID_TYPE_SDID, dstRankInfo.superDeviceId, remoteLinkInfo.ip));
      91              :             } else {
      92            0 :                 CHK_RET(hrtRaGetSingleSocketVnicIpInfo(
      93              :                     devicePhyId_, DeviceIdType::DEVICE_ID_TYPE_PHY_ID, dstRankInfo.devicePhyId, remoteLinkInfo.ip));
      94              :             }
      95              :             // 通信域未分配端口则使用默认端口
      96              :             remoteLinkInfo.port
      97           10 :                 = dstRankInfo.deviceVnicPort == HCCL_INVALID_PORT ? HETEROG_CCL_PORT : dstRankInfo.deviceVnicPort;
      98           10 :             remoteLinkInfo.socketsPerLink = 1;
      99           10 :             string newTag = GenerateSocketTag(devicePhyId_, rankInfoList_[i].devicePhyId);
     100           10 :             std::vector<std::shared_ptr<HcclSocket>> tmpSockets;
     101              :             HcclResult ret
     102           10 :                 = socketManager_->CreateSingleLinkSocket(newTag, vnicPortCtx_, remoteLinkInfo, tmpSockets, false, true);
     103           10 :             CHK_PRT_RET(
     104              :                 ret != HCCL_SUCCESS,
     105              :                 HCCL_ERROR(
     106              :                     "[Create][DestSockets]Create single link sockets failed, "
     107              :                     "local rank[%u], remote rank[%u]",
     108              :                     userRank_, rankInfoList_[i].userRank),
     109              :                 ret);
     110           10 :             if (tmpSockets.size() != 1) {
     111            0 :                 HCCL_ERROR(
     112              :                     "[SymmetricMemoryAgent][CreateVnic] socket number[%llu] is not 1 as expected!", tmpSockets.size());
     113            0 :                 return HCCL_E_INTERNAL;
     114              :             }
     115              :             // 设置强制断链为关闭,避免进程退出时recv失败
     116           10 :             tmpSockets[0]->SetForceClose(false);
     117           10 :             mapRankIdconnectedSockets_[remoteLinkInfo.userRank] = (tmpSockets[0]);
     118           10 :             mapRankId2DevPhyId_[remoteLinkInfo.userRank] = remoteLinkInfo.devicePhyId;
     119           10 :         }
     120              :     }
     121              : 
     122           18 :     for (const auto& kv : mapRankIdconnectedSockets_) {
     123           10 :         CHK_PRT_RET(
     124              :             socketManager_->WaitLinkEstablish(kv.second) != HCCL_SUCCESS,
     125              :             HCCL_ERROR(
     126              :                 "[SymmetricMemoryAgent][EstablishSockets] tag[%s] socket establish failed",
     127              :                 kv.second->GetTag().c_str()),
     128              :             HCCL_E_INTERNAL);
     129              :     }
     130            8 :     return HCCL_SUCCESS;
     131              : }
     132              : 
     133           10 : std::string SymmetricMemoryAgent::GenerateSocketTag(u32 localRank, u32 remoteRank)
     134              : {
     135           10 :     u32 small = localRank;
     136           10 :     u32 large = remoteRank;
     137              : 
     138           10 :     if (localRank > remoteRank) {
     139            0 :         small = remoteRank;
     140            0 :         large = localRank;
     141              :     }
     142              : 
     143              :     // Socket构造规则:前缀 + identifier + small + large
     144              :     std::string tag
     145           10 :         = STR_IPC_MEM_EXCHANGE + "_" + identifier_ + "_" + std::to_string(small) + ":" + std::to_string(large);
     146           10 :     return tag;
     147              : }
     148              : 
     149           10 : HcclResult SymmetricMemoryAgent::ExchangeInfo(void* inputPtr, void* outputPtr, u64 inputSize)
     150              : {
     151           10 :     CHK_PTR_NULL(inputPtr);
     152            9 :     CHK_PTR_NULL(outputPtr);
     153            8 :     CHK_PRT_RET(inputSize == 0, HCCL_ERROR("Input size is 0"), HCCL_E_PARA);
     154              :     // 校验 inputSize 是否超过协议载荷上限
     155            7 :     CHK_PRT_RET(
     156              :         inputSize > PACKET_DATA_MAX_LEN,
     157              :         HCCL_ERROR("Input size %lu exceeds max payload %u", inputSize, PACKET_DATA_MAX_LEN), HCCL_E_PARA);
     158              :     // 校验是否建链成功
     159            6 :     CHK_PRT_RET(
     160              :         mapRankIdconnectedSockets_.find(rightRank_) == mapRankIdconnectedSockets_.end(),
     161              :         HCCL_ERROR("[ExchangeInfo] rightRank_%u socket not found in map", rightRank_), HCCL_E_INTERNAL);
     162            4 :     CHK_PRT_RET(
     163              :         mapRankIdconnectedSockets_.find(leftRank_) == mapRankIdconnectedSockets_.end(),
     164              :         HCCL_ERROR("[ExchangeInfo] leftRank_%u socket not found in map", leftRank_), HCCL_E_INTERNAL);
     165              : 
     166            4 :     HCCL_INFO(
     167              :         "[SymmetricMemoryAgent] start to ExchangeInfo, inputPtr[%p], outputPtr[%p], inputSize[%llu]", inputPtr,
     168              :         outputPtr, inputSize);
     169              : 
     170              :     // 重置本轮状态
     171            4 :     outputDataPtr_ = static_cast<u8*>(outputPtr);
     172            4 :     currentInputSize_ = inputSize; // 记录实际有效长度
     173            4 :     collectedCount_ = 0;
     174              :     // 本地数据处理:先把自己的一份拷到 Output 对应位置
     175            4 :     u8* selfDstPtr = outputDataPtr_ + (userRank_ * inputSize);
     176            4 :     CHK_SAFETY_FUNC_RET(memcpy_s(selfDstPtr, inputSize, inputPtr, inputSize));
     177            4 :     collectedCount_++;
     178              : 
     179              :     Packet dataPkt;
     180            4 :     dataPkt.type = MsgType::MSG_TYPE_DATA;
     181            4 :     dataPkt.rankId = userRank_;
     182            4 :     CHK_SAFETY_FUNC_RET(memset_s(dataPkt.data, PACKET_DATA_MAX_LEN, 0, PACKET_DATA_MAX_LEN));
     183            4 :     CHK_SAFETY_FUNC_RET(memcpy_s(dataPkt.data, PACKET_DATA_MAX_LEN, inputPtr, inputSize));
     184              :     {
     185            4 :         std::lock_guard<std::mutex> lock(queueMutex_);
     186            4 :         requestQueue_.push(dataPkt);
     187            4 :     }
     188            4 :     isProcessingTask_ = true;
     189              : 
     190            4 :     CHK_RET(WaitForCollectionComplete());
     191            1 :     HCCL_INFO("[SymmetricMemoryAgent] ExchangeInfo end");
     192            1 :     return HCCL_SUCCESS;
     193              : }
     194              : 
     195            4 : HcclResult SymmetricMemoryAgent::WaitForCollectionComplete()
     196              : {
     197            4 :     std::unique_lock<std::mutex> lock(completionMutex_);
     198            4 :     auto timeout = std::chrono::seconds(GetExternalInputHcclLinkTimeOut());
     199            4 :     auto status = completionCv_.wait_for(lock, timeout);
     200            4 :     if (status == std::cv_status::timeout) {
     201            6 :         HCCL_ERROR("[SymmetricMemoryAgent] ExchangeInfo Timeout! Collected: %u/%u", collectedCount_.load(), rankSize_);
     202            3 :         return HCCL_E_TCP_TRANSFER;
     203              :     }
     204            1 :     return HCCL_SUCCESS;
     205            4 : }
     206              : 
     207            9 : void SymmetricMemoryAgent::DealWithRequest()
     208              : {
     209            9 :     if (hrtSetDevice(deviceLogicId_) != HCCL_SUCCESS) {
     210            0 :         return;
     211              :     }
     212              : 
     213            9 :     std::vector<u8> leftRecvBuf(PACKET_TOTAL_LEN, 0);
     214            9 :     u32 leftRecvLen = 0;
     215              : 
     216         1899 :     while (threadRun_) {
     217         1890 :         if (isProcessingTask_) {
     218         1883 :             if (collectedCount_ < rankSize_) {
     219          944 :                 u64 received = 0;
     220          944 :                 std::unique_lock<std::mutex> lock(socketMutex_);
     221          944 :                 HcclResult ret = mapRankIdconnectedSockets_[leftRank_]->IRecv(
     222          944 :                     leftRecvBuf.data() + leftRecvLen, PACKET_TOTAL_LEN - leftRecvLen, received);
     223              : 
     224          944 :                 CHK_PRT_CONT(
     225              :                     (ret != HCCL_SUCCESS) && (ret != HCCL_E_AGAIN),
     226              :                     HCCL_ERROR(
     227              :                         "[SymmetricMemoryAgent][DealWithRequest] IRecv failed, ret[%d] remoteRank[%u] "
     228              :                         "receivedSize[%llu]",
     229              :                         ret, leftRank_, leftRecvLen));
     230              : 
     231          944 :                 leftRecvLen += received;
     232          944 :                 if (leftRecvLen == PACKET_TOTAL_LEN) {
     233            4 :                     Packet* pkt = reinterpret_cast<Packet*>(leftRecvBuf.data());
     234            4 :                     ProcessReceivedPacket(*pkt);
     235            4 :                     leftRecvLen = 0;
     236              :                 }
     237          944 :             }
     238         1883 :             std::lock_guard<std::mutex> lock(queueMutex_);
     239         1883 :             if (!requestQueue_.empty()) {
     240          943 :                 Packet pkt = requestQueue_.front();
     241          943 :                 std::unique_lock<std::mutex> sockLock(socketMutex_);
     242              :                 HcclResult ret
     243          943 :                     = mapRankIdconnectedSockets_[rightRank_]->Send(static_cast<void*>(&pkt), PACKET_TOTAL_LEN);
     244          943 :                 if (ret == HCCL_SUCCESS) {
     245            2 :                     requestQueue_.pop();
     246              :                 } else {
     247          941 :                     HCCL_ERROR(
     248              :                         "[SymmetricMemoryAgent][DealWithRequest] Data(from rank[%u]) Send to rank[%u] failed.",
     249              :                         pkt.rankId, rightRank_);
     250              :                 }
     251          943 :             }
     252              :             // 检查是否完全结束, 退出条件: 数据全齐 && 队列空闲
     253         1883 :             if (requestQueue_.empty() && collectedCount_ == rankSize_) {
     254            1 :                 std::unique_lock<std::mutex> lock(completionMutex_);
     255            1 :                 HCCL_INFO("[SymmetricMemoryAgent] ExchangeInfo Complete.");
     256            1 :                 isProcessingTask_ = false;
     257            1 :                 completionCv_.notify_all();
     258            1 :             }
     259         1883 :         }
     260         1890 :         SaluSleep(USLEEP_ONE_THOUSAND);
     261              :     }
     262              : 
     263            9 :     hrtResetDevice(deviceLogicId_);
     264            9 : }
     265              : 
     266            4 : HcclResult SymmetricMemoryAgent::ProcessReceivedPacket(Packet& pkt)
     267              : {
     268            4 :     if (pkt.rankId < rankSize_ && pkt.rankId != userRank_) {
     269            4 :         u8* dest = outputDataPtr_ + (pkt.rankId * currentInputSize_);
     270            4 :         CHK_SAFETY_FUNC_RET(memcpy_s(dest, currentInputSize_, pkt.data, currentInputSize_));
     271            4 :         collectedCount_++;
     272              :     }
     273            8 :     HCCL_INFO(
     274              :         "[SymmetricMemoryAgent][ProcessReceivedPacket] Data Recv from rank[%u]. Collected[%u / %u].", pkt.rankId,
     275              :         collectedCount_.load(), rankSize_);
     276              :     // Ring 转发逻辑:如果数据不是自己的,也不是右边Rank发出的(转了一圈),则转发给右边
     277            4 :     if (pkt.rankId != userRank_ && pkt.rankId != rightRank_) {
     278            0 :         std::lock_guard<std::mutex> lock(queueMutex_);
     279            0 :         requestQueue_.push(pkt);
     280            0 :     }
     281            4 :     return HCCL_SUCCESS;
     282              : }
     283              : } // namespace hccl
        

Generated by: LCOV version 2.0-1