LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_ccu_context/all_to_all - ccu_context_all_to_all_mesh1d.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 92 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 3 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 "ccu_context_all_to_all_mesh1d.h"
      12              : #include "ccu_instruction_all_to_all_mesh1d.h"
      13              : 
      14              : namespace Hccl {
      15              : 
      16              : constexpr int CKE_IDX_0 = 0;
      17              : constexpr int CKE_IDX_1 = 1;
      18              : constexpr int CKE_IDX_2 = 2;
      19              : 
      20            0 : CcuContextAllToAllMesh1D::CcuContextAllToAllMesh1D(
      21            0 :     const CcuCtxArg& arg, const std::vector<CcuTransport*>& transports, const CcuTransportGroup& group)
      22            0 :     : CcuContextAlgBase(arg, transports, group)
      23              : {
      24            0 :     const CcuCtxArgAllToAllMesh1D* ctxArg = dynamic_cast<const CcuCtxArgAllToAllMesh1D*>(&arg);
      25            0 :     if (ctxArg == nullptr) {
      26            0 :         THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D::ctxArg ptr is null"));
      27              :     }
      28            0 :     rankId_ = ctxArg->rankId;
      29            0 :     if (ctxArg->dimSize.size() > 0) {
      30            0 :         rankSize_ = ctxArg->dimSize[0];
      31              :     }
      32            0 :     if (transports.size() == 0) {
      33            0 :         THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D transports is empty"));
      34              :     }
      35            0 :     loadFromMem_ = ctxArg->loadFromMem;
      36            0 : }
      37              : 
      38            0 : void CcuContextAllToAllMesh1D::Algorithm()
      39              : {
      40            0 :     HCCL_INFO("[ccuAllToAllMesh1D_context] AllToAllMesh1D run.");
      41              :     // 创建Variable,用于交换地址及token
      42            0 :     u32 transportId = 0;
      43            0 :     for (u64 id = 0; id < rankSize_; id++) {
      44            0 :         if (id == rankId_) {
      45            0 :             input_.push_back(CreateVariable());
      46            0 :             output_.push_back(CreateVariable());
      47            0 :             token_.push_back(CreateVariable());
      48              :         } else { // 非本地,使用远端Variable
      49            0 :             CHK_PRT_RET(
      50              :                 transports[transportId] == nullptr,
      51              :                 HCCL_ERROR("[CcuContextAllToAllMesh1D] Algorithm transport ptr is null"), );
      52            0 :             input_.push_back(CreateVariable((*transports[transportId]), CKE_IDX_0));
      53            0 :             output_.push_back(CreateVariable((*transports[transportId]), CKE_IDX_1));
      54            0 :             token_.push_back(CreateVariable((*transports[transportId]), CKE_IDX_2));
      55            0 :             transportId++;
      56              :         }
      57              :     }
      58            0 :     sliceSize_ = CreateVariable();
      59            0 :     srcStride_ = CreateVariable();
      60            0 :     srcOffset_ = CreateVariable();
      61            0 :     dstOffset_ = CreateVariable();
      62            0 :     groupOpSize_ = CreateGroupOpSize();
      63              : 
      64              :     // 从SQE load args,本rank需要的input、output地址等信息
      65              :     // inputAddr, outputAddr, tokenInfo, srcStride, srcOffset, dstOffset, groupOpSize
      66            0 :     Load(input_[rankId_]);
      67            0 :     Load(output_[rankId_]);
      68            0 :     Load(token_[rankId_]);
      69            0 :     Load(sliceSize_); // 本轮传输的分片大小
      70            0 :     Load(srcStride_); // 单片数据大小
      71            0 :     Load(srcOffset_);
      72            0 :     Load(dstOffset_);
      73            0 :     Load(groupOpSize_);
      74              : 
      75              :     // 前同步。交换信息,将本Rank load的in\out等地址信息写到所有对端的对应Variable中,并同步
      76            0 :     uint16_t selfBit = 1 << rankId_; // 本rank的mask
      77            0 :     uint16_t allBit = ((1 << rankSize_) - 1) & (~(1 << rankId_));
      78              : 
      79            0 :     srcOffset_ += input_[rankId_];
      80              : 
      81            0 :     for (auto t : transports) {
      82              :         // (transport, param, paramID, SemID, mask)
      83            0 :         WriteVariableWithSignal(*t, output_[rankId_], CKE_IDX_1, CKE_IDX_1, selfBit); // index = 1,传递output信息
      84            0 :         WriteVariableWithSignal(*t, token_[rankId_], CKE_IDX_2, CKE_IDX_2, selfBit);  // index = 2,传递token信息
      85              :     }
      86              : 
      87            0 :     GroupWait(*transportGroup, CKE_IDX_1, allBit); // index = 1,传递output信息
      88            0 :     GroupWait(*transportGroup, CKE_IDX_2, allBit); // index = 2,传递token信息
      89              : 
      90              :     // 创建GSA, src为本地的各片HBM地址GSA列表,dst为所有对端的HBM地址GSA列表
      91            0 :     std::vector<CcuRep::Memory> src;
      92            0 :     for (uint64_t rankIdx = 0; rankIdx < rankSize_; rankIdx++) {
      93            0 :         src.push_back(CreateMemory());
      94              :     }
      95            0 :     std::vector<CcuRep::Memory> dst;
      96            0 :     for (uint64_t rankIdx = 0; rankIdx < rankSize_; rankIdx++) {
      97            0 :         dst.push_back(CreateMemory());
      98              :     }
      99              : 
     100              :     // 考虑stride信息
     101            0 :     for (uint64_t r = 0; r < rankSize_; r++) {
     102            0 :         src[r].token = token_[r];
     103            0 :         dst[r].token = token_[r];
     104              : 
     105              :         // src[r] = srcOffset + r*srcStride
     106            0 :         src[r].addr = srcOffset_;
     107            0 :         for (uint64_t i = 0; i < r; i++) {
     108            0 :             src[r].addr += srcStride_;
     109              :         }
     110              :         // dst[r] = recvBuf[r] + dstOffset
     111            0 :         dst[r].addr = output_[r];
     112            0 :         dst[r].addr += dstOffset_;
     113              :     }
     114              : 
     115              :     // 创建CKE,源端保序
     116            0 :     CcuRep::MaskSignal locMask = CreateMaskSignal();
     117              :     //  all2all 数据搬运
     118            0 :     transportId = 0;
     119            0 :     for (uint64_t r = 0; r < rankSize_; r++) {
     120            0 :         if (r != rankId_) {
     121            0 :             Write(*transports[transportId], dst[r], src[r], sliceSize_, locMask, 1 << r);
     122            0 :             transportId++;
     123              :         }
     124              :     }
     125            0 :     GroupCopy(dst[rankId_], src[rankId_], groupOpSize_);
     126            0 :     LocalWait(locMask, allBit);
     127              : 
     128              :     //  后同步
     129            0 :     for (auto t : transports) {
     130            0 :         if (t == nullptr) {
     131            0 :             THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D::Algorithm transport ptr is null"));
     132              :         }
     133            0 :         RemotePost(*t, CKE_IDX_0, selfBit);
     134              :     }
     135            0 :     GroupWait(*transportGroup, CKE_IDX_0, allBit);
     136            0 :     HCCL_INFO("[AllToAllAlgo] AllToAllMesh1D end");
     137              : 
     138            0 :     return;
     139            0 : }
     140              : 
     141            0 : std::vector<uint64_t> CcuContextAllToAllMesh1D::GeneArgs(const CcuTaskArg& arg)
     142              : {
     143            0 :     const CcuTaskArgAllToAllMesh1D* taskArg = dynamic_cast<const CcuTaskArgAllToAllMesh1D*>(&arg);
     144            0 :     if (taskArg == nullptr) {
     145            0 :         THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D::taskArg ptr is null"));
     146              :     }
     147            0 :     uint64_t inputAddr = taskArg->inputAddr;
     148            0 :     uint64_t outputAddr = taskArg->outputAddr;
     149            0 :     uint64_t tokenInfo = taskArg->token;
     150              : 
     151            0 :     uint64_t srcStride = taskArg->srcStride;
     152            0 :     uint64_t srcOffset = taskArg->srcOffset;
     153            0 :     uint64_t dstOffset = taskArg->dstOffset;
     154              : 
     155            0 :     uint64_t sliceSize = taskArg->sliceSize;
     156            0 :     auto goSize = CalGoSize(sliceSize);
     157            0 :     HCCL_INFO(
     158              :         "[AllToAllAlgo] inputAddr[%llu], outputAddr[%llu], sliceSize[%llu], srcStride[%llu], srcOffset[%llu], "
     159              :         "dstOffset[%llu].",
     160              :         inputAddr, outputAddr, sliceSize, srcStride, srcOffset, dstOffset);
     161              : 
     162              :     return {inputAddr, outputAddr, tokenInfo, sliceSize, srcStride, srcOffset,
     163            0 :             dstOffset, goSize[0],  goSize[1], goSize[2], goSize[3]};
     164            0 : }
     165              : 
     166              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1