Fichiers
rocm-systems/src/include/group.h
T

117 lignes
3.7 KiB
C
Brut Vue normale Historique

2018-09-24 16:06:59 -07:00
/*************************************************************************
* Copyright (c) 2015-2017, NVIDIA CORPORATION. All rights reserved.
*
* See LICENSE.txt for license information
************************************************************************/
#ifndef NCCL_GROUP_H_
#define NCCL_GROUP_H_
#include "nccl.h"
2019-11-19 14:57:39 -08:00
#include "comm.h"
2018-09-24 16:06:59 -07:00
2022-05-24 02:02:31 -07:00
ncclResult_t ncclGroupErrCheck(ncclResult_t ret);
void ncclGroupCommJoin(struct ncclComm* comm);
void ncclGroupCommPreconnect(struct ncclComm* comm);
2022-08-18 02:53:17 -07:00
ncclResult_t ncclGroupCommLeave(struct ncclComm* comm);
void ncclGroupJobAbort();
2018-09-24 16:06:59 -07:00
2019-11-19 14:57:39 -08:00
typedef ncclResult_t(*ncclInitFunc_t)(ncclComm_t* newcomm, int ndev, ncclUniqueId commId, int myrank, int cudaDev);
2018-09-24 16:06:59 -07:00
2019-11-19 14:57:39 -08:00
ncclResult_t ncclAsyncInit(ncclInitFunc_t func, ncclComm_t* newcomm, int ndev, ncclUniqueId commId, int myrank, int cudaDev);
2018-09-24 16:06:59 -07:00
2022-08-18 02:53:17 -07:00
typedef enum ncclGroupJobState {
ncclGroupJobRunning = 0,
ncclGroupJobDone = 1,
ncclGroupJobJoined = 2,
} ncclGroupJobState_t;
2022-05-24 02:02:31 -07:00
struct ncclAsyncJob {
struct ncclAsyncJob* next;
pthread_t thread;
ncclResult_t result;
ncclResult_t(*func)(struct ncclAsyncJob*);
void(*undo)(struct ncclAsyncJob*);
void(*destructor)(void*);
2022-08-18 02:53:17 -07:00
ncclGroupJobState_t state;
volatile uint32_t *abortFlag; /* point to comm abortFlag */
ncclComm_t comm;
2022-05-24 02:02:31 -07:00
};
ncclResult_t ncclAsyncLaunch(
struct ncclAsyncJob* job,
ncclResult_t(*func)(struct ncclAsyncJob*),
void(*undo)(struct ncclAsyncJob*),
2022-08-18 02:53:17 -07:00
void(*destructor)(void*), ncclComm_t comm
2022-05-24 02:02:31 -07:00
);
2022-08-18 02:53:17 -07:00
struct ncclGroupJob {
struct ncclAsyncJob base;
struct ncclComm **groupCommHeadPtr;
struct ncclComm **groupCommPreconnectHeadPtr;
ncclResult_t *groupErrorPtr;
volatile bool *abortFlagPtr;
struct ncclIntruQueue<struct ncclAsyncJob, &ncclAsyncJob::next> *asyncJobsPtr;
bool doneFlag;
};
2022-05-24 02:02:31 -07:00
ncclResult_t ncclGroupStartInternal();
ncclResult_t ncclGroupEndInternal();
2022-08-18 02:53:17 -07:00
ncclResult_t ncclAsyncJobComplete(struct ncclAsyncJob* job);
2022-05-24 02:02:31 -07:00
////////////////////////////////////////////////////////////////////////////////
extern __thread int ncclGroupDepth; // depth of ncclGroupStart nesting
extern __thread ncclResult_t ncclGroupError;
extern __thread struct ncclComm* ncclGroupCommHead;
extern __thread struct ncclComm* ncclGroupCommPreconnectHead;
2022-08-18 02:53:17 -07:00
extern __thread int ncclGroupBlocking;
2022-05-24 02:02:31 -07:00
inline ncclResult_t ncclGroupStartInternal() {
ncclGroupDepth++;
return ncclSuccess;
}
inline ncclResult_t ncclGroupErrCheck(ncclResult_t ret) {
if (ncclGroupDepth > 0) {
2022-08-18 02:53:17 -07:00
if (ret != ncclSuccess && ret != ncclInProgress) ncclGroupError = ret;
2022-05-24 02:02:31 -07:00
}
return ret;
}
// Add comm to this thread's group
inline void ncclGroupCommJoin(struct ncclComm* comm) {
if (comm->groupNext == reinterpret_cast<struct ncclComm*>(0x1)) {
// Insert comm into ncclGroupCommHead adjacent to sibling comms. This preserves
// the users program order yet insures siblings occur consecutively. This
// is required by doLaunches() in "group.cc".
struct ncclComm** pp = &ncclGroupCommHead;
while (*pp != nullptr && comm->intraComm0 != (*pp)->intraComm0)
pp = &(*pp)->groupNext;
comm->groupNext = *pp;
*pp = comm;
// Comms gets a new memory stack scope upon joining. Each task batched for
// this comm is allocated there.
ncclMemoryStackPush(&comm->memScoped);
}
2022-08-18 02:53:17 -07:00
ncclGroupBlocking = comm->blocking;
2022-05-24 02:02:31 -07:00
}
// Add comm to this thread's group needing preconnect
inline void ncclGroupCommPreconnect(struct ncclComm* comm) {
if (comm->preconnectNext == reinterpret_cast<struct ncclComm*>(0x1)) {
comm->preconnectNext = ncclGroupCommPreconnectHead;
ncclGroupCommPreconnectHead = comm;
}
}
// Comm has left group
2022-08-18 02:53:17 -07:00
inline ncclResult_t ncclGroupCommLeave(struct ncclComm* comm) {
2022-05-24 02:02:31 -07:00
comm->groupNext = reinterpret_cast<struct ncclComm*>(0x1);
ncclMemoryStackPop(&comm->memScoped);
2022-08-18 02:53:17 -07:00
return ncclSuccess;
2022-05-24 02:02:31 -07:00
}
2018-09-24 16:06:59 -07:00
#endif