LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_ccu_context/all_to_all - ccu_context_half_alltoallv_mesh1d.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 169 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 8 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 <random>
      12              : #include <algorithm>
      13              : #include "ccu_context_half_alltoallv_mesh1d.h"
      14              : #include "ccu_instruction_half_alltoallv_mesh1d.h"
      15              : 
      16              : namespace Hccl {
      17              : 
      18              : constexpr uint16_t RANK_NUM_PER_CKE = 16; // 本rank给远端置位时应当写的CKE,16个对端一个CKE
      19              : 
      20            0 : void CcuContextHalfAllToAllVMesh1D::ExchangeCtxResource()
      21              : {
      22            0 :     if (missionId_ == 1) {
      23            0 :         ExportVariable(userInAddr_, ctxName_ + "_UserInAddr_" + std::to_string(missionId_));
      24            0 :         ExportVariable(sendSizeAddr_, ctxName_ + "_SendSizeAddr_" + std::to_string(missionId_));
      25            0 :         ExportVariable(sendOffsetAddr_, ctxName_ + "_SendOffsetAddr__" + std::to_string(missionId_));
      26            0 :         ExportVariable(recvOffset_, ctxName_ + "_RecvOffset_" + std::to_string(missionId_));
      27              :     } else {
      28            0 :         anoUserInAddr_ = ImportVariable(ctxName_ + "_UserInAddr_" + std::to_string(1 - missionId_));
      29            0 :         anoSendSizeAddr_ = ImportVariable(ctxName_ + "_SendSizeAddr_" + std::to_string(1 - missionId_));
      30            0 :         anoSendOffsetAddr_ = ImportVariable(ctxName_ + "_SendOffsetAddr__" + std::to_string(1 - missionId_));
      31            0 :         anoRecvOffset_ = ImportVariable(ctxName_ + "_RecvOffset_" + std::to_string(1 - missionId_));
      32              :     }
      33            0 :     ExportMaskSignal(locMiSignal0_, ctxName_ + "_MiSync0_" + std::to_string(missionId_));
      34            0 :     anoMiSignal0_ = ImportMaskSignal(ctxName_ + "_MiSync0_" + std::to_string(1 - missionId_));
      35            0 :     ExportMaskSignal(locMiSignal1_, ctxName_ + "_MiSync1_" + std::to_string(missionId_));
      36            0 :     anoMiSignal1_ = ImportMaskSignal(ctxName_ + "_MiSync1_" + std::to_string(1 - missionId_));
      37              : 
      38            0 :     return;
      39              : }
      40              : 
      41            0 : CcuContextHalfAllToAllVMesh1D::CcuContextHalfAllToAllVMesh1D(
      42            0 :     const CcuCtxArg& arg, const std::vector<CcuTransport*>& transports, const CcuTransportGroup& group)
      43            0 :     : CcuContextAlgBase(arg, transports, group)
      44              : {
      45            0 :     const CcuCtxArgHalfAllToAllVMesh1D* ctxArg = dynamic_cast<const CcuCtxArgHalfAllToAllVMesh1D*>(&arg);
      46            0 :     if (ctxArg == nullptr) {
      47            0 :         THROW<NullPtrException>(StringFormat("CcuContextHalfAllToAllVMesh1D::ctxArg ptr is null"));
      48              :     }
      49            0 :     ctxName_ = ctxArg->GetCtxSignature().Describe();
      50            0 :     rankId_ = ctxArg->rankId;
      51            0 :     if (ctxArg->dimSize.size() > 0) {
      52            0 :         rankSize_ = ctxArg->dimSize[0];
      53              :     }
      54            0 :     missionId_ = ctxArg->missionId;
      55            0 :     signalNum_ = (rankSize_ + RANK_NUM_PER_CKE - 1) / RANK_NUM_PER_CKE; // 每个CKE有16个bit
      56            0 :     myCclBufferAddr_ = ctxArg->cclBufferAddr;
      57              : 
      58            0 :     userInAddr_ = CreateVariable();
      59            0 :     sendSizeAddr_ = CreateVariable();
      60            0 :     sendOffsetAddr_ = CreateVariable();
      61            0 :     recvOffset_ = CreateVariable();
      62            0 :     goSize_ = CreateGroupOpSize();
      63              : 
      64            0 :     curSrc_ = CreateMemory();
      65            0 :     if (transports.size() == 0 || transports.size() < rankSize_ - 1) {
      66            0 :         THROW<NullPtrException>(StringFormat("CcuContextHalfAllToAllVMesh1D transports is empty or size is less"));
      67              :     }
      68            0 :     for (uint32_t i = 0; i < rankSize_; i++) {
      69              :         // curDst在初始化时即赋值各个rank的cclBuffer地址
      70            0 :         if (i == rankId_) {
      71            0 :             curDst_.emplace_back(CreateMemory());
      72              :         } else {
      73            0 :             uint16_t transIdx = (i < rankId_) ? i : i - 1;
      74            0 :             curDst_.emplace_back(GetRmtBuffer(*transports[transIdx], 0));
      75              :         }
      76            0 :         token_.emplace_back(CreateVariable());
      77            0 :         sendSizeA_.emplace_back(CreateVariable());
      78            0 :         sendSizeB_.emplace_back(CreateVariable());
      79            0 :         sendOffsetA_.emplace_back(CreateVariable());
      80            0 :         sendOffsetB_.emplace_back(CreateVariable());
      81              :     }
      82            0 :     ccuStartSignal_ = CreateMaskSignal();
      83            0 :     ccuEndSignal_ = CreateMaskSignal();
      84            0 :     for (uint32_t i = 0; i < signalNum_; i++) {
      85            0 :         writeDoneSignal_.emplace_back(CreateMaskSignal());
      86              :     }
      87              : 
      88              :     // mission间交互的资源
      89            0 :     locMiSignal0_ = CreateMaskSignal();
      90            0 :     locMiSignal1_ = CreateMaskSignal();
      91            0 :     ExchangeCtxResource();
      92              : 
      93            0 :     return;
      94            0 : }
      95              : 
      96            0 : void CcuContextHalfAllToAllVMesh1D::LoadArgs()
      97              : {
      98            0 :     if (missionId_ == 0) {
      99            0 :         Load(userInAddr_);
     100            0 :         Load(sendSizeAddr_);
     101            0 :         Load(token_[rankId_]);
     102            0 :         Load(sendOffsetAddr_);
     103            0 :         Load(goSize_);
     104            0 :         Load(recvOffset_);
     105              : 
     106              :         // Mi0将参数同步给Mi1
     107            0 :         LocalCtxPostVar(userInAddr_, anoUserInAddr_, anoMiSignal0_, 1 << 0);         // 用第1个bit标记
     108            0 :         LocalCtxPostVar(sendSizeAddr_, anoSendSizeAddr_, anoMiSignal0_, 1 << 1);     // 用第2个
     109            0 :         LocalCtxPostVar(sendOffsetAddr_, anoSendOffsetAddr_, anoMiSignal0_, 1 << 2); // 用第3个
     110            0 :         LocalCtxPostVar(recvOffset_, anoRecvOffset_, anoMiSignal0_, 1 << 3);         // 用第4个
     111              :     } else {
     112            0 :         LocalWait(locMiSignal0_, (1 << 4) - 1); // 共同步4个参数
     113              :     }
     114            0 :     return;
     115              : }
     116              : 
     117            0 : void CcuContextHalfAllToAllVMesh1D::LoadArgsFromMem()
     118              : {
     119              :     // 暂不支持用单条指令加载多个参数
     120            0 :     CcuRep::Variable dataLength = CreateVariable();
     121            0 :     CcuRep::Variable tempAddr = CreateVariable();
     122            0 :     LoadArgs();
     123              : 
     124              :     // 加载本端的cclbuffer地址
     125            0 :     curDst_[rankId_].addr = myCclBufferAddr_;
     126              : 
     127              :     // 连续加载rankSize * 2个sendSize
     128            0 :     dataLength = 8; // 每个Xn占8个byte
     129              : 
     130            0 :     u32 argsCount = sendSizeA_.size() + sendSizeB_.size() + sendOffsetA_.size() + sendOffsetB_.size();
     131            0 :     std::vector<CcuRep::Variable> tempArgs(argsCount);
     132            0 :     HCCL_INFO("CcuContextHalfAllToAllVMesh1D LoadArgsFromMem, argsCount:[%u]", argsCount);
     133              : 
     134            0 :     for (uint32_t i = 0; i < tempArgs.size(); ++i) {
     135            0 :         tempArgs[i] = CreateContinuousVariable();
     136              :     }
     137            0 :     LoadVariable(sendSizeAddr_, tempArgs[0], argsCount);
     138              : 
     139            0 :     u32 argIdx = 0;
     140            0 :     for (uint32_t i = 0; i < sendSizeA_.size(); i++) {
     141            0 :         sendSizeA_[i] = tempArgs[argIdx];
     142            0 :         argIdx++;
     143              :     }
     144            0 :     for (uint32_t i = 0; i < sendSizeB_.size(); i++) {
     145            0 :         sendSizeB_[i] = tempArgs[argIdx];
     146            0 :         argIdx++;
     147              :     }
     148            0 :     for (uint32_t i = 0; i < sendOffsetA_.size(); i++) {
     149            0 :         sendOffsetA_[i] = tempArgs[argIdx];
     150            0 :         argIdx++;
     151              :     }
     152            0 :     for (uint32_t i = 0; i < sendOffsetB_.size(); i++) {
     153            0 :         sendOffsetB_[i] = tempArgs[argIdx];
     154            0 :         argIdx++;
     155              :     }
     156            0 :     return;
     157            0 : }
     158              : 
     159            0 : void CcuContextHalfAllToAllVMesh1D::MissionSync(uint32_t signalIndex)
     160              : {
     161            0 :     const uint32_t MISSION_NUM = 2;
     162            0 :     if (signalIndex > 1) {
     163            0 :         THROW<InvalidParamsException>(
     164            0 :             StringFormat("[CcuContextHalfAllToAllVMesh1D] Unexpected SignalInex[%u]", signalIndex));
     165              :     }
     166            0 :     LocalCtxPost(anoMiSignal1_, 1 << (missionId_ + signalIndex * MISSION_NUM));
     167            0 :     LocalWait(locMiSignal1_, 1 << (1 - missionId_ + signalIndex * MISSION_NUM));
     168            0 :     return;
     169              : }
     170              : 
     171            0 : void CcuContextHalfAllToAllVMesh1D::PostSync()
     172              : {
     173            0 :     if (missionId_ == 0) {
     174            0 :         uint16_t signalId = rankId_ / RANK_NUM_PER_CKE;
     175            0 :         uint16_t selfBit = 1 << (rankId_ % RANK_NUM_PER_CKE);
     176            0 :         for (auto t : transports) {
     177            0 :             if (t == nullptr) {
     178            0 :                 THROW<NullPtrException>(StringFormat("CcuContextHalfAllToAllVMesh1D::Algorithm transport ptr is null"));
     179              :             }
     180            0 :             RemotePost(*t, signalId, selfBit);
     181              :         }
     182              : 
     183            0 :         for (uint16_t sId = 0; sId < signalNum_; sId++) {
     184              :             uint32_t waitBit;
     185            0 :             if (sId != signalNum_ - 1) {
     186            0 :                 waitBit = (1 << RANK_NUM_PER_CKE) - 1; // 等待全部16个peer
     187              :             } else {
     188            0 :                 waitBit = ((1 << (rankSize_ - (signalNum_ - 1) * RANK_NUM_PER_CKE)) - 1);
     189              :             }
     190            0 :             if (sId == signalId) {
     191            0 :                 waitBit &= ~selfBit; // 如果这个CKE上有自己对应的bit,设为0
     192              :             }
     193            0 :             GroupWait(*transportGroup, sId, waitBit);
     194              :         }
     195              :     }
     196            0 :     return;
     197              : }
     198              : 
     199            0 : void CcuContextHalfAllToAllVMesh1D::Algorithm()
     200              : {
     201            0 :     HCCL_INFO("[CcuContextHalfAllToAllVMesh1D] AllgatherMesh1D Algorithm Begins.");
     202            0 :     LoadArgsFromMem();
     203              : 
     204              :     // 向每个对端发送数据
     205            0 :     CcuRep::Memory lgSrc = CreateMemory();
     206            0 :     CcuRep::Memory lgDst = CreateMemory();
     207            0 :     CcuRep::Variable tempCount = CreateVariable();
     208            0 :     lgSrc.token = token_[rankId_];
     209            0 :     lgDst.token = token_[rankId_];
     210            0 :     curSrc_.token = token_[rankId_];
     211              : 
     212            0 :     for (uint32_t peerId = 0; peerId < rankSize_; peerId++) {
     213            0 :         CcuRep::Variable& curCount = missionId_ == 0 ? sendSizeA_[peerId] : sendSizeB_[peerId];
     214            0 :         CcuRep::Variable& curOffset = missionId_ == 0 ? sendOffsetA_[peerId] : sendOffsetB_[peerId];
     215            0 :         uint16_t peerSignalId = peerId / RANK_NUM_PER_CKE;
     216            0 :         uint16_t peerBit = 1 << (peerId % RANK_NUM_PER_CKE);
     217              : 
     218            0 :         tempCount = curCount;
     219            0 :         curSrc_.addr = userInAddr_;
     220            0 :         curSrc_.addr += curOffset;
     221            0 :         curDst_[peerId].addr += recvOffset_;
     222            0 :         if (missionId_ == 1) {
     223            0 :             curDst_[peerId].addr += sendSizeA_[peerId]; // Mi1的dst需要加chunkOffset
     224              :         }
     225            0 :         if (peerId == rankId_) {
     226            0 :             lgSrc.addr = curSrc_.addr;
     227            0 :             lgDst.addr = curDst_[peerId].addr;
     228            0 :             LocalPost(writeDoneSignal_[peerSignalId], peerBit);
     229              :         } else {
     230            0 :             uint16_t transIdx = (peerId < rankId_) ? peerId : peerId - 1;
     231            0 :             CCU_IF(tempCount != 0)
     232              :             {
     233            0 :                 Write(
     234            0 :                     *(transports[transIdx]), curDst_[peerId], curSrc_, tempCount, writeDoneSignal_[peerSignalId],
     235              :                     peerBit);
     236            0 :             }
     237            0 :             CCU_IF(tempCount == 0) { LocalPost(writeDoneSignal_[peerSignalId], peerBit); }
     238              :         }
     239              :     }
     240              : 
     241            0 :     if (missionId_ == 0) {
     242            0 :         GroupCopy(lgDst, lgSrc, goSize_);
     243              :     }
     244            0 :     for (uint16_t sId = 0; sId < signalNum_; sId++) {
     245              :         uint32_t waitBit;
     246            0 :         if (sId != signalNum_ - 1) {
     247            0 :             waitBit = (1 << RANK_NUM_PER_CKE) - 1; // 等待全部16个peer
     248              :         } else {
     249            0 :             waitBit = ((1 << (rankSize_ - (signalNum_ - 1) * RANK_NUM_PER_CKE)) - 1);
     250              :         }
     251            0 :         LocalWait(writeDoneSignal_[sId], waitBit);
     252              :     }
     253              : 
     254            0 :     MissionSync(0);
     255            0 :     PostSync();
     256            0 :     HCCL_INFO("[CcuContextHalfA2AUnions] Algorithm Ends.");
     257            0 :     return;
     258            0 : }
     259              : 
     260            0 : std::vector<uint64_t> CcuContextHalfAllToAllVMesh1D::GeneArgs(const CcuTaskArg& arg)
     261              : {
     262              :     (void)arg;
     263            0 :     std::vector<uint64_t> args = {};
     264            0 :     uint64_t argNum = missionId_ == 0 ? 9 : 0; // Mi0有9个Load,Mi1不做Load
     265            0 :     for (uint32_t i = 0; i < argNum; i++) {
     266            0 :         args.emplace_back(0);
     267              :     }
     268            0 :     return args;
     269            0 : }
     270              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1