LCOV - code coverage report
Current view: top level - legacy/ascend910/framework/group - hccl_group.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 69.7 % 198 138
Test Date: 2026-07-28 12:11:00 Functions: 87.5 % 16 14

            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 <algorithm>
      12              : #include <mutex>
      13              : #include <thread>
      14              : #include <vector>
      15              : 
      16              : #include "hccl_group.h"
      17              : 
      18              : using namespace hccl;
      19              : 
      20              : thread_local s32 hcclGroupDepth = 0;
      21              : thread_local std::deque<std::shared_ptr<struct hcclAsyncJob>> hcclInitJobs;
      22              : thread_local std::vector<HcclComm> hcclGroupCommList;
      23              : 
      24            2 : HcclResult HcclLegacyGroupStart()
      25              : {
      26            2 :     hcclGroupDepth++;
      27            2 :     HCCL_INFO("[HcclGroupStart] hcclGroupDepth=[%d]", hcclGroupDepth);
      28            2 :     return HCCL_SUCCESS;
      29              : }
      30              : 
      31              : namespace hccl{
      32              : 
      33            0 : HcclResult initGroupPlanner(HcclComm comm) {
      34            0 :     hccl::hcclComm* hcclComm = static_cast<hccl::hcclComm *>(comm);
      35            0 :     std::shared_ptr<struct hcclKernelPlanner> planner = hcclComm->planner;
      36            0 :     u32 rankSize = INVALID_VALUE_RANKSIZE;
      37            0 :     hcclComm->GetRankSize(rankSize);
      38            0 :     planner->rankSize = rankSize;
      39            0 :     HCCL_DEBUG("[initGroupPlanner] ranksize: %d", rankSize);
      40              :     
      41            0 :     planner->nTasksP2p = 0;
      42            0 :     planner->nTasksColl = 0;
      43            0 :     return HCCL_SUCCESS;
      44            0 : }
      45              : 
      46            3 : HcclResult taskAppend(HcclComm comm, hcclOpInfo& info) {
      47            3 :     hccl::hcclComm* hcclComm = static_cast<hccl::hcclComm *>(comm);
      48            3 :     std::shared_ptr<struct hcclKernelPlanner> planner = hcclComm->planner;
      49            3 :     if(planner->nTasksP2p == -1){
      50            0 :         initGroupPlanner(comm);
      51              :     }
      52              : 
      53            3 :     HcclResult ret = HCCL_SUCCESS;
      54            3 :     if (info.coll == HcclCMDType::HCCL_CMD_SEND || info.coll == HcclCMDType::HCCL_CMD_RECEIVE) {
      55            2 :         hcclComm->SetGroupMode(true);
      56            2 :         bool isSendOp = (info.coll == HcclCMDType::HCCL_CMD_SEND);
      57              :         HcclSendRecvItem item;
      58            2 :         item.sendRecvType = isSendOp ? HcclSendRecvType::HCCL_SEND : HcclSendRecvType::HCCL_RECV;
      59            2 :         item.buf = isSendOp ? info.sendbuff : info.recvbuff;
      60            2 :         item.count  = isSendOp ? info.sendCount : info.recvCount;
      61            2 :         item.dataType = isSendOp ? info.sendType : info.recvType;
      62            2 :         item.remoteRank = info.root;
      63              : 
      64            2 :         planner->sendRecvInfo.push_back(item);
      65              : 
      66            2 :         if (planner->sendRecvMainStream == nullptr) { // 用第一条用户流作为主流
      67            2 :             planner->sendRecvMainStream = info.stream;
      68            2 :             HCCL_INFO("[TaskAppend] planner->sendRecvMainStream[%p]", planner->sendRecvMainStream);
      69              :         }
      70              :         
      71            2 :         planner->nTasksP2p += 1;
      72            2 :     }
      73              :     else {
      74            1 :         hcclOpInfo task = info;
      75            1 :         planner->collTaskQueue.push_back(task);
      76            1 :         planner->nTasksColl += 1;
      77              :         /*记录stream到planner*/
      78            1 :         planner->collStreams.insert(info.stream);
      79              :     }
      80              : 
      81            3 :     auto itComm = std::find(hcclGroupCommList.begin(), hcclGroupCommList.end(), comm);
      82            3 :     if (itComm == hcclGroupCommList.end()) {
      83            3 :         hcclGroupCommList.push_back(comm);
      84              :     }
      85            3 :     return ret;
      86            3 : }
      87              : 
      88            2 : HcclResult commInitTaskAppend(std::shared_ptr<struct hcclAsyncJob> job, HcclResult (*func)(struct hcclAsyncJob *), HcclComm* comm)
      89              : {
      90            2 :     HCCL_INFO("[hcclAsyncJobEnqueue] add item to queue");
      91            2 :     CHK_PRT_RET(!job, HCCL_ERROR("[commInitTaskAppend] job is nullptr"), HCCL_E_INTERNAL);
      92              :     /*hcclAsyncLaunch只是将job放入队列,并不等待执行完成。groupLaunch->asyncJobLaunch中给每个job起一个线程去执行*/
      93            1 :     job->func = func;
      94            1 :     job->comm = comm;
      95            1 :     job->state = hcclGroupJobRunning;
      96            1 :     hcclInitJobs.push_back(job);
      97            1 :     return HCCL_SUCCESS;
      98              : }
      99              : }// namespace hccl
     100              : 
     101            1 : void *hcclAsyncJobMain(void *arg)
     102              : {
     103            1 :     struct hcclAsyncJob *job = (struct hcclAsyncJob *)arg;
     104            1 :     job->result = job->func(job); /*func是上层asyncjob里面设置的函数*/
     105            1 :     if (job->result == HCCL_SUCCESS) {
     106            1 :         HCCL_INFO("Function launch success");
     107              :     }
     108              :     /*加锁修改job->state为hcclGroupJobDone*/
     109            1 :     std::unique_lock<std::mutex> lock(job->mtx);
     110            1 :     job->state = hcclGroupJobDone;
     111            1 :     return arg;
     112            1 : }
     113              : 
     114            1 : static HcclResult asyncJobLaunch()
     115              : {
     116            1 :     HCCL_DEBUG("[asyncJobLaunch] entered");
     117            1 :     HcclResult ret = HCCL_SUCCESS;
     118            1 :     bool jobsDone = false;
     119              : 
     120            1 :     if (!hcclInitJobs.empty()) {
     121            2 :         for (auto job : hcclInitJobs) {
     122            1 :             CHK_PRT_RET(!job, HCCL_ERROR("[asyncJobLaunch] job is nullptr"), HCCL_E_INTERNAL);
     123            1 :             job->thread.reset(new (std::nothrow) std::thread(&hcclAsyncJobMain, job.get()));
     124            1 :             CHK_PRT_RET(!job->thread, HCCL_ERROR("[asyncJobLaunch]threads reset failed "), HCCL_E_INTERNAL);
     125            1 :         }
     126              : 
     127              :         do { /*主线程轮询阻塞,等待所有线程上的asyncJob执行完成*/
     128            1 :             jobsDone = true;
     129            2 :             for (auto job : hcclInitJobs) {
     130              :                 /*上面job执行线程可能并发修改state,在主线程里面要通过加线程锁来读取*/
     131            1 :                 hcclGroupJobState_t state = hcclGroupJobJoined;
     132            1 :                 std::unique_lock<std::mutex> lock(job->mtx);
     133            1 :                 state = job->state;
     134              : 
     135            1 :                 if (state == hcclGroupJobRunning) {
     136            0 :                     jobsDone = false;
     137            1 :                 } else if (state == hcclGroupJobDone) {
     138            1 :                     job->thread->join();
     139            1 :                     job->state = hcclGroupJobJoined;
     140            1 :                     if (job->result != HCCL_SUCCESS && ret == HCCL_SUCCESS) {
     141            0 :                         ret = job->result;
     142              :                     }
     143              :                 } else {
     144              :                     /* safety check */
     145            0 :                     CHK_PRT_RET(state != hcclGroupJobJoined,
     146              :                         HCCL_ERROR("[asyncJobLaunch] state != hcclGroupJobJoined"),
     147              :                         HCCL_E_INTERNAL);
     148              :                 }
     149            1 :             }
     150              :             // Let preconnect threads progress.
     151            1 :             if (jobsDone == false)
     152            0 :                 usleep(1);
     153            1 :         } while (jobsDone == false);
     154            1 :         hcclInitJobs.clear();
     155            1 :         CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[asyncJobLaunch] fail!"), ret);
     156              :     }
     157            1 :     return HCCL_SUCCESS;
     158              : }
     159              : 
     160              : // 计算集合通信算子的数据量(字节),用于按数据量降序排列(仅非V类算子参与排序)
     161         1034 : u64 calcOpDataVolume(const hcclOpInfo& info)
     162              : {
     163         1034 :     switch (info.coll) {
     164         1012 :         case HcclCMDType::HCCL_CMD_ALLGATHER:
     165              :         case HcclCMDType::HCCL_CMD_ALLTOALL:
     166              :         case HcclCMDType::HCCL_CMD_ALLREDUCE:
     167              :         case HcclCMDType::HCCL_CMD_BROADCAST:
     168              :         case HcclCMDType::HCCL_CMD_REDUCE:
     169         1012 :             return info.sendCount * SIZE_TABLE[info.sendType];
     170           15 :         case HcclCMDType::HCCL_CMD_REDUCE_SCATTER:
     171              :         case HcclCMDType::HCCL_CMD_SCATTER:
     172           15 :             return info.recvCount * SIZE_TABLE[info.recvType];
     173            7 :         default:
     174            7 :             return 0;
     175              :     }
     176              : }
     177              : 
     178              : // 对group任务排序:非V类按数据量降序在前,V类保持原序在后
     179           11 : std::vector<hcclOpInfo> sortGroupTasks(const std::deque<hcclOpInfo>& tasks)
     180              : {
     181           11 :     std::vector<hcclOpInfo> sorted(tasks.begin(), tasks.end());
     182          129 :     auto isVariant = [](const hcclOpInfo& op) {
     183          253 :         return op.coll == HcclCMDType::HCCL_CMD_ALLTOALLV ||
     184          124 :                op.coll == HcclCMDType::HCCL_CMD_ALLTOALLVC ||
     185          370 :                op.coll == HcclCMDType::HCCL_CMD_ALLGATHER_V ||
     186          246 :                op.coll == HcclCMDType::HCCL_CMD_REDUCE_SCATTER_V;
     187              :     };
     188           11 :     auto itV = std::stable_partition(sorted.begin(), sorted.end(),
     189          129 :         [&isVariant](const hcclOpInfo& op) { return !isVariant(op); });
     190           11 :     std::stable_sort(sorted.begin(), itV,
     191          398 :         [](const hcclOpInfo& a, const hcclOpInfo& b) {
     192          398 :             return calcOpDataVolume(a) > calcOpDataVolume(b);
     193              :         });
     194           11 :     return sorted;
     195            0 : }
     196              : 
     197            1 : static HcclResult doLaunches(HcclComm comm)
     198              : {
     199            1 :     hccl::hcclComm* hcclComm = static_cast<hccl::hcclComm *>(comm);
     200            1 :     std::shared_ptr<struct hcclKernelPlanner> planner = hcclComm->planner;
     201            1 :     HcclUs startutime = TIME_NOW();
     202            1 :     if (planner->nTasksP2p != 0) { 
     203              :         // 将所有send/recv的任务打包作为一个集合通信算子来执行
     204            1 :         HCCL_INFO("HcclBatchSendRecvGroup, sendRecvInfo.size()[%u]", static_cast<u32>(planner->sendRecvInfo.size()));
     205            1 :         CHK_RET(HcclBatchSendRecvGroup(planner->sendRecvInfo.data(), planner->sendRecvInfo.size(), comm, planner->sendRecvMainStream));
     206              :     }
     207            1 :     HCCL_INFO("[doLaunches] take time [%lld]us.", DURATION_US(TIME_NOW() - startutime));
     208            1 :     if (planner->nTasksColl != 0) {
     209              :         /* 按数据量降序排列,V类算子各rank数据量可能不同,统一放到最后,保持原序 */
     210            1 :         auto sortedTasks = sortGroupTasks(planner->collTaskQueue);
     211            1 :         HCCL_INFO("Collectives sorted!");
     212            1 :         for (auto& taskColl : sortedTasks) {
     213            0 :             switch (taskColl.coll) {
     214            0 :                 case HcclCMDType::HCCL_CMD_ALLGATHER:
     215            0 :                     HCCL_INFO("AllGather, sendCount[%llu]", taskColl.sendCount);
     216            0 :                     CHK_RET(HcclAllGatherInner(taskColl.sendbuff, taskColl.recvbuff, taskColl.sendCount,
     217              :                                     taskColl.sendType, taskColl.comm, taskColl.stream));
     218            0 :                     break;
     219            0 :                 case HcclCMDType::HCCL_CMD_REDUCE_SCATTER:
     220            0 :                     HCCL_INFO("ReduceScatter, recvCount[%llu]", taskColl.recvCount);
     221            0 :                     CHK_RET(HcclReduceScatterInner(taskColl.sendbuff, taskColl.recvbuff, taskColl.recvCount, 
     222              :                                 taskColl.recvType, taskColl.op, taskColl.comm, taskColl.stream));
     223            0 :                     break;
     224            0 :                 case HcclCMDType::HCCL_CMD_ALLREDUCE:
     225            0 :                     HCCL_INFO("AllReduce, sendCount[%llu]", taskColl.sendCount);
     226            0 :                     CHK_RET(HcclAllReduceInner(taskColl.sendbuff, taskColl.recvbuff, taskColl.sendCount, taskColl.sendType, 
     227              :                                     taskColl.op, taskColl.comm, taskColl.stream));
     228            0 :                     break;
     229            0 :                 case HcclCMDType::HCCL_CMD_BROADCAST:
     230            0 :                     CHK_RET(HcclBroadcastInner(taskColl.sendbuff, taskColl.sendCount, taskColl.sendType, taskColl.root, taskColl.comm, 
     231              :                                     taskColl.stream));
     232            0 :                     break;
     233            0 :                 case HcclCMDType::HCCL_CMD_ALLTOALL:
     234            0 :                     CHK_RET(HcclAlltoAllInner(taskColl.sendbuff, taskColl.sendCount, taskColl.sendType, taskColl.recvbuff, taskColl.recvCount, taskColl.recvType, 
     235              :                                     taskColl.comm, taskColl.stream));
     236            0 :                     break;
     237            0 :                 case HcclCMDType::HCCL_CMD_ALLTOALLV:
     238            0 :                     CHK_RET(HcclAlltoAllVInner(taskColl.sendbuff, taskColl.sendCounts, taskColl.sdispls, taskColl.sendType,
     239              :                                     taskColl.recvbuff, taskColl.recvCounts, taskColl.rdispls, taskColl.recvType,
     240              :                                     taskColl.comm, taskColl.stream));
     241            0 :                     break;
     242            0 :                 case HcclCMDType::HCCL_CMD_ALLTOALLVC:
     243            0 :                     CHK_RET(HcclAlltoAllVCInner(taskColl.sendbuff, taskColl.sendCounts, taskColl.sendType, taskColl.recvbuff, taskColl.recvType, 
     244              :                                     taskColl.comm, taskColl.stream));
     245            0 :                     break;
     246            0 :                 case HcclCMDType::HCCL_CMD_REDUCE:
     247            0 :                     CHK_RET(HcclReduceInner(taskColl.sendbuff, taskColl.recvbuff, taskColl.recvCount, taskColl.recvType, taskColl.op, taskColl.root, 
     248              :                                     taskColl.comm, taskColl.stream));
     249            0 :                     break;
     250            0 :                 case HcclCMDType::HCCL_CMD_SCATTER:
     251            0 :                     CHK_RET(HcclScatterInner(taskColl.sendbuff, taskColl.recvbuff, taskColl.recvCount, taskColl.recvType, taskColl.root, taskColl.comm, taskColl.stream));
     252            0 :                     break;
     253            0 :                 case HcclCMDType::HCCL_CMD_ALLGATHER_V:
     254            0 :                     CHK_RET(HcclAllGatherVInner(taskColl.sendbuff, taskColl.sendCount, taskColl.recvbuff, taskColl.recvCounts, taskColl.rdispls, 
     255              :                                     taskColl.recvType, taskColl.comm, taskColl.stream));
     256            0 :                     break;
     257            0 :                 case HcclCMDType::HCCL_CMD_REDUCE_SCATTER_V:
     258            0 :                     CHK_RET(HcclReduceScatterVInner(taskColl.sendbuff, taskColl.sendCounts, taskColl.sdispls, taskColl.recvbuff, taskColl.recvCount, 
     259              :                                         taskColl.sendType, taskColl.op, taskColl.comm, taskColl.stream));
     260            0 :                     break;
     261            0 :                 default:
     262            0 :                     HCCL_ERROR("[doLaunches] not supported hcclFunc!");
     263            0 :                     return HCCL_E_INTERNAL;
     264              :             }
     265              :         }
     266            1 :     }
     267            1 :     return HCCL_SUCCESS;
     268            1 : }
     269              : 
     270            1 : static HcclResult groupLaunch()
     271              : {  // 将各种通信域初始化/destroy的asyncJobs,在这里触发放到背景线程执行
     272            1 :     HCCL_INFO("[groupLaunch] entered");
     273              : 
     274            1 :     asyncJobLaunch();
     275            1 :     HCCL_DEBUG("[groupLaunch] asyncJobLaunch done");
     276            2 :     for (HcclComm comm : hcclGroupCommList) {
     277            1 :         doLaunches(comm);
     278              :     }
     279            1 :     HCCL_INFO("[groupLaunch] doLaunches done");
     280              :     //流同步
     281            2 :     for (HcclComm comm : hcclGroupCommList){
     282            1 :         hccl::hcclComm* hcclComm = static_cast<hccl::hcclComm *>(comm);
     283            1 :         std::shared_ptr<struct hcclKernelPlanner> planner = hcclComm->planner;
     284            1 :         for (auto it : planner->collStreams){
     285            0 :             CHK_RET(hcclStreamSynchronize(it));
     286              :         }
     287            1 :     }
     288            1 :     HCCL_INFO("groupLauch Done!");
     289            1 :     return HCCL_SUCCESS;
     290              : }
     291              : 
     292            1 : inline void groupLocalResetJobState()
     293              : {
     294              :     // hcclcomm中group相关的变量
     295            2 :     for (HcclComm comm : hcclGroupCommList) {
     296            1 :         hccl::hcclComm* hcclComm = static_cast<hccl::hcclComm *>(comm);
     297            1 :         hcclComm->planner = std::make_shared<hcclKernelPlanner>();
     298            1 :         hcclComm->SetGroupMode(false);
     299              :     }
     300            1 :     hcclGroupCommList.clear();
     301              : 
     302            1 :     return;
     303              : }
     304              : 
     305            1 : HcclResult HcclLegacyGroupEnd()
     306              : {
     307            1 :     groupLaunch();
     308            1 :     HCCL_INFO("[GroupEnd] done groupLaunch");
     309            1 :     groupLocalResetJobState();
     310            1 :     HCCL_INFO("[GroupEnd] to the end");
     311            1 :     return HCCL_SUCCESS;
     312              : }
     313              : 
     314            0 : HcclResult HcclLegacyAsyncJobLaunch()
     315              : {
     316            0 :     return asyncJobLaunch();
     317              : }
        

Generated by: LCOV version 2.0-1