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

Generated by: LCOV version 2.0-1