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

Generated by: LCOV version 2.0-1