Merge remote-tracking branch 'nccl/master' into develop

This commit is contained in:
BertanDogancay
2025-06-20 07:53:59 -05:00
136 changed files with 8510 additions and 5421 deletions
+50 -12
View File
@@ -141,7 +141,7 @@ namespace {
}
#endif
// Final wait/copy.
prims.directRecv(offset, offset, nelem);
prims.directRecv(offset, nelem);
#if defined(ENABLE_NPKIT) && defined(ENABLE_NPKIT_EVENT_ALL_GATHER_RING_DIRECT_RECV_EXIT)
if (tid == 0) {
@@ -220,25 +220,63 @@ struct RunWorkColl<ncclFuncAllGather, T, RedOp, NCCL_ALGO_RING, NCCL_PROTO_LL128
template<typename T, typename RedOp>
struct RunWorkColl<ncclFuncAllGather, T, RedOp, NCCL_ALGO_PAT, NCCL_PROTO_SIMPLE> {
__device__ __forceinline__ void run(int tid, int nthreads, struct ncclDevWorkColl* work) {
#if __CUDA_ARCH__ >= 600
using Proto = ProtoSimple<1, 1>;
const int nranks = ncclShmem.comm.nRanks;
const int rank = ncclShmem.comm.rank;
size_t count, channelOffset, channelCount, chunkCount;
ncclCollCbdPart(work, ncclShmem.channelId, Proto::Id, sizeof(T), &count, &channelOffset, &channelCount, &chunkCount);
T *inputBuf = (T*)work->sendbuff;
T *outputBuf = (T*)work->recvbuff;
Primitives<T, RedOp, FanSymmetric<1>, 0, Proto, 0> prims
(tid, nthreads, NULL, NULL, inputBuf, outputBuf, work->redOpArg, 0*Proto::MaxGroupWidth, 0, 0, nullptr, nullptr, 0, primsModePatAg);
static constexpr int nworkers = NCCL_PAT_NWORKERS;
struct ncclPatShmem* shmem = (struct ncclPatShmem*)ncclScratchForWarp(0);
uint64_t pollCount = 0;
__syncthreads(); // Don't start using shared mem until everyone arrives
for (int i=tid; i<NCCL_SHMEM_PAT_STEPS; i+=nthreads) shmem->patSteps[i].flags = 0;
if (tid == 0) shmem->localAccSize = 0;
if (tid == nworkers) shmem->parallelFactor = 0;
__syncthreads();
PatAGAlgorithm<T> patAlgo(chunkCount*sizeof(T), NCCL_STEPS, channelOffset, channelOffset + channelCount, count, chunkCount, rank, nranks);
int last = 0;
while (!last) {
int recvDim, sendDim, recvOffset, sendOffset, recvStepOffset, postRecv, postSend, nelem;
size_t inpIx, outIx;
patAlgo.getNextOp(recvDim, sendDim, inpIx, outIx, recvOffset, sendOffset, recvStepOffset, nelem, postRecv, postSend, last);
prims.patCopy(recvDim, sendDim, inpIx, outIx, recvOffset, sendOffset, recvStepOffset, nelem, postRecv, postSend);
if (tid == nworkers) { // Algo computation thread
PatAGAlgorithm<T> patAlgo(chunkCount*sizeof(T), NCCL_STEPS, NCCL_PAT_NWORKERS/WARP_SIZE, channelOffset, channelOffset + channelCount, count, chunkCount, rank, nranks);
int parallelFactor = shmem->parallelFactor = patAlgo.getParallelFactor();
int step = 0;
while (1) {
struct ncclPatStep* ps = shmem->patSteps+(step%NCCL_SHMEM_PAT_STEPS);
cuda::atomic_ref<int, cuda::thread_scope_block> poll(ps->flags);
while (poll.load(cuda::memory_order_acquire) != 0) pollCount++; // Wait for workers to be done with step 'step-NCCL_SHMEM_PAT_STEPS'
patAlgo.getNextOp(ps);
int last = ps->last;
step++;
if (last == 2) break;
}
} else if (tid < nworkers) { // Worker threads
T *inputBuf = (T*)work->sendbuff;
T *outputBuf = (T*)work->recvbuff;
int parallelFactor = 0;
volatile int* pfPtr = &shmem->parallelFactor;
while (parallelFactor == 0) parallelFactor = *pfPtr;
int groupSize = nworkers/(WARP_SIZE*parallelFactor) * WARP_SIZE;
int group = tid / groupSize;
int nGroups = nworkers / groupSize;
int tidInGroup = tid - group*groupSize;
// We don't use recvPeers/sendPeers so let's pass shmem structs instead
Primitives<T, RedOp, FanSymmetric<1>, 0, Proto, 0> prims
(tidInGroup, groupSize, (int*)shmem->recvDims, (int*)shmem->sendDims, inputBuf, outputBuf, work->redOpArg, group, 0, 0, nullptr, nullptr, 0, primsModePatAg);
int step = group;
while(1) {
struct ncclPatStep* ps = shmem->patSteps+(step%NCCL_SHMEM_PAT_STEPS);
cuda::atomic_ref<int, cuda::thread_scope_block> poll(ps->flags);
while (poll.load(cuda::memory_order_acquire) == 0) pollCount++; // Wait for compute thread
int last = ps->last;
prims.patCopy(ps, shmem);
if (tidInGroup == 0) poll.store(0, cuda::memory_order_release); // Return element to compute thread
if (last) break;
step += nGroups;
}
}
#endif
}
};
+5 -5
View File
@@ -190,7 +190,7 @@ namespace {
offset = gridOffset + elemOffset + chunkOffset;
nelem = (int)min(chunkCount, remCount - chunkOffset);
prims.directRecv(offset, offset, nelem);
prims.directRecv(offset, nelem);
#if defined(ENABLE_NPKIT) && defined(ENABLE_NPKIT_EVENT_ALL_REDUCE_RING_DIRECT_RECV_EXIT)
if (tid == 0) {
@@ -329,7 +329,7 @@ namespace {
for (size_t elemOffset = 0; elemOffset < channelCount; elemOffset += chunkCount) {
offset = gridOffset + elemOffset;
nelem = min(chunkCount, channelCount - elemOffset);
prims.directRecv(offset, offset, nelem);
prims.directRecv(offset, nelem);
}
}
else {
@@ -528,7 +528,7 @@ namespace {
for (size_t elemOffset = 0; elemOffset < channelCount; elemOffset += chunkCount) {
offset = gridOffset + elemOffset;
nelem = min(chunkCount, channelCount - elemOffset);
prims.directRecv(offset, offset, nelem);
prims.directRecv(offset, nelem);
}
}
else {
@@ -1055,7 +1055,7 @@ struct RunWorkColl<ncclFuncAllReduce, T, RedOp, NCCL_ALGO_COLLNET_CHAIN, NCCL_PR
for (ssize_t gridOffset = 0; gridOffset < size; gridOffset += loopSize) {
ssize_t offset = gridOffset + bid * int(chunkSize);
int nelem = min(chunkSize, size - offset);
prims.directRecv(offset, offset, nelem, /*postOp*/true);
prims.directRecv(offset, nelem, /*postOp*/true);
}
}
} else {
@@ -1082,7 +1082,7 @@ struct RunWorkColl<ncclFuncAllReduce, T, RedOp, NCCL_ALGO_COLLNET_CHAIN, NCCL_PR
for (ssize_t gridOffset = 0; gridOffset < size; gridOffset += loopSize) {
ssize_t offset = gridOffset + bid*int(chunkSize);
int nelem = min(chunkSize, size-offset);
prims.directRecv(offset, offset, nelem);
prims.directRecv(offset, nelem);
}
} else {
for (ssize_t gridOffset = 0; gridOffset < size; gridOffset += loopSize) {
+1 -1
View File
@@ -83,7 +83,7 @@ namespace {
prims.directCopySend(offset, offset, nelem);
}
} else if (nextRank == root) {
prims.directRecv(offset, offset, nelem);
prims.directRecv(offset, nelem);
} else {
prims.directRecvCopyDirectSend(offset, offset, nelem);
}
+54 -30
View File
@@ -144,6 +144,8 @@ struct ncclShmemData {
int nWorks;
int workSize;
uint32_t workConsumed;
uint64_t workCounter;
bool profilerEnabled;
struct ncclShmemGroup groups[NCCL_MAX_GROUPS];
uint64_t redOpArgs[NCCL_MAX_NVLS_ARITY+1];
@@ -236,24 +238,6 @@ __device__ inline bool barrier_red_or(bool vote, int name, int nThreads) {
: "=r"(ans) : "r"((int)vote), "r"(name), "r"(nThreads) : "memory");
return bool(ans);
}
__device__ inline bool barrier_red_or_aligned(bool vote, int name) {
int ans;
asm volatile("{ .reg .pred p;"
" setp.ne.s32 p, %1, 0;"
" barrier.red.or.pred.aligned p, %2, p; "
" selp.s32 %0, 1, 0, p; }"
: "=r"(ans) : "r"((int)vote), "r"(name) : "memory");
return bool(ans);
}
__device__ inline bool barrier_red_or_aligned(bool vote, int name, int nThreads) {
int ans;
asm("{ .reg .pred p;"
" setp.ne.s32 p, %1, 0;"
" barrier.red.or.pred.aligned p, %2, %3, p; "
" selp.s32 %0, 1, 0, p; }"
: "=r"(ans) : "r"((int)vote), "r"(name), "r"(nThreads) : "memory");
return bool(ans);
}
#ifdef ENABLE_PROFILING
#define __insert_timestamp(line_num) do { \
@@ -455,6 +439,48 @@ struct RunWorkBatch {
}
};
#define START 0
#define STOP 1
#define FINI 2
__device__ __forceinline__ bool profilerEnabled(void) {
// Check if any of the workItems in the batch is profiled. If so, there is an equivalent
// profiler ProxyOp waiting for the counter update in the host thread. If this check was
// done only for the first workItem the profiler counter for other workItems in the batch
// could never be updated, leaving the host thread spinning forever for the counter update
// and causing a hang.
bool enabled = false;
for (int i = 0; i < ncclShmem.nWorks && !enabled; i++) {
if (ncclShmem.workType == ncclDevWorkTypeP2p)
enabled = ((struct ncclDevWorkP2p*)ncclShmem.workStorage)[i].profilerEnabled;
else
enabled = ((struct ncclDevWorkColl*)ncclShmem.workStorage)[i].profilerEnabled;
}
return enabled;
}
__device__ __forceinline__ void profiler(int action) {
if (action == START) {
if (threadIdx.x == 0) {
// increment workCounter regardless of the profiler being active or not
ncclShmem.channel.workCounter += ncclShmem.nWorks;
if(!profilerEnabled()) return;
ncclShmem.comm.workStarted[ncclShmem.channelId] = ncclShmem.channel.workCounter;
}
} else if (action == STOP) {
if (threadIdx.x == 0 && profilerEnabled()) {
ncclShmem.comm.workCompleted[ncclShmem.channelId] = ncclShmem.channel.workCounter;
}
} else { // FINI
if (threadIdx.x == 0) {
// store the workCounter back to vidmem regardless of the profiler being active or not
((ncclDevCommAndChannels*)ncclShmem.args.comm)->channels[ncclShmem.channelId].workCounter = ncclShmem.channel.workCounter;
if (!profilerEnabled()) return;
ncclShmem.comm.workCompleted[ncclShmem.channelId] = ncclShmem.channel.workCounter;
}
}
}
template<int SpecializedFnId, typename SpecializedRunWorkBatch, bool COLLTRACE, int COLL_UNROLL>
__device__ __forceinline__ void ncclKernelMain(struct ncclDevKernelArgs const* args) {
const int tid = threadIdx.x;
@@ -517,8 +543,13 @@ __device__ __forceinline__ void ncclKernelMain(struct ncclDevKernelArgs const* a
break;
}
__syncthreads(); // publish ncclShmem.{args, channelId}
/* set abort flag to 0 */
if (tid == 0) {
ncclShmem.aborted = 0;
ncclShmem.channel.workCounter = ((ncclDevCommAndChannels*)ncclShmem.args.comm)->channels[ncclShmem.channelId].workCounter;
}
// Use first 2 warps to load comm and channel, and reamaining load work batch.
// Use first 2 warps to load comm and channel, and remaining load work batch.
switch (tid/WARP_SIZE) {
case 0:
{ void* dst = &ncclShmem.comm;
@@ -566,9 +597,9 @@ __device__ __forceinline__ void ncclKernelMain(struct ncclDevKernelArgs const* a
ncclShmem.comm.workConsumed[ncclShmem.channelId] = ncclShmem.workConsumed;
}
while (true) {
while (ncclShmem.aborted == 0) {
if (tid == 0) __insert_timestamp(__LINE__);
profiler(START);
if (0 <= SpecializedFnId && ncclShmem.funcId == (unsigned)SpecializedFnId) {
SpecializedRunWorkBatch().run();
} else {
@@ -586,21 +617,14 @@ __device__ __forceinline__ void ncclKernelMain(struct ncclDevKernelArgs const* a
default:
break;
}
profiler(STOP);
loadWorkBatchToShmem(tid%WARP_SIZE, tn, args, batchIx);
__syncthreads();
// Check whether the last operation was aborted and make sure all threads exit
bool aborted = false;
if (tid == 0) aborted = *ncclShmem.comm.abortFlag;
aborted = __any(aborted); // publish ncclShmem.work
if (tid == 0 && ncclShmem.args.workStorageType == ncclDevWorkStorageTypeFifo) {
// ncclShmem.workConsumed written by loadWorkBatchToShmem before barrier_red_or()
// ncclShmem.workConsumed written by loadWorkBatchToShmem before __syncthreads()
ncclShmem.comm.workConsumed[ncclShmem.channelId] = ncclShmem.workConsumed;
}
if (aborted) {
if(COLLTRACE && tid%WARP_SIZE == 0) traceAbort();
break;
}
if (COLLTRACE && tid%WARP_SIZE == 0) traceKernelLaunch(ncclCollTraceCollLaunchType, batchIx);
}
if (COLLTRACE && tid%WARP_SIZE == 0) traceKernelEnd(ncclCollTraceKernelEndType);
+14 -2
View File
@@ -13,7 +13,7 @@
#include "common_kernel.h"
#include "common.h"
#define NCCL_SPINS_BEFORE_CHECK_ABORT 1000000
#define NCCL_SPINS_BEFORE_CHECK_ABORT 10000
#define barrier_by_group_common(__THREAD_FENCE) do { \
if (nthreads == NCCL_MAX_NTHREADS) { \
@@ -154,7 +154,7 @@ struct PrimitivesWithoutDirect {
__device__ void directSendFromOutput(intptr_t outIx, int eltN) {
static_cast<RealPrimitives*>(this)->sendFromOutput(outIx, eltN);
}
__device__ void directRecv(intptr_t inpIx, intptr_t outIx, int eltN) {
__device__ void directRecv(intptr_t outIx, int eltN) {
static_cast<RealPrimitives*>(this)->recv(outIx, eltN, /*postOp=*/false);
}
__device__ void directCopySend(intptr_t inpIx, intptr_t outIx, int eltN, bool postOp=false) {
@@ -178,6 +178,18 @@ struct PrimitivesWithoutDirect {
}
};
__device__ inline int checkAbort(int &abortCache, const int abortValue, int &spins) {
if (abortCache & abortValue) return 1;
if (++spins < NCCL_SPINS_BEFORE_CHECK_ABORT) return 0;
spins = 0;
int abort = __atomic_load_n((ncclShmem.comm.abortFlag), __ATOMIC_SEQ_CST);
if (abort) {
__atomic_store_n(&ncclShmem.aborted, abort, __ATOMIC_SEQ_CST);
abortCache |= abortValue;
}
return abort;
}
#include "prims_simple.h"
#include "prims_ll.h"
#include "prims_ll128.h"
+13 -10
View File
@@ -85,15 +85,18 @@ private:
#endif
}
uint32_t abort = 0;
int abort = 0;
inline __device__ int checkAbort(int &spins, int send) {
spins++;
if (abort == 0 && spins == NCCL_SPINS_BEFORE_CHECK_ABORT) {
abort = __atomic_load_n((ncclShmem.comm.abortFlag), __ATOMIC_SEQ_CST);
__device__ inline int checkAbort(int &abortCache, const int abortValue, int &spins) {
if (abortCache == 0 && ++spins == NCCL_SPINS_BEFORE_CHECK_ABORT) {
int abort = __atomic_load_n((ncclShmem.comm.abortFlag), __ATOMIC_SEQ_CST);
spins = 0;
if (abort) {
__atomic_store_n(&ncclShmem.aborted, abort, __ATOMIC_SEQ_CST);
abortCache |= abortValue;
}
}
return abort;
return abortCache;
}
inline __device__ void waitSend(int nbytes) {
@@ -108,7 +111,7 @@ private:
while (sendConnHeadCache + NCCL_STEPS < sendConnHead + 1) {
__builtin_amdgcn_s_sleep(1);
sendConnHeadCache = atomicAdd((unsigned long long *)sendConnHeadPtr, 0);
if (checkAbort(spins, 1)) break;
if (checkAbort(abort, 1, spins)) break;
}
if (sendConnFifo) {
int size = ((sendConnHead & NCCL_LL_CLEAN_MASK) == NCCL_LL_CLEAN_MASK) ? stepLines*sizeof(union ncclLLFifoLine) : nbytes;
@@ -168,7 +171,7 @@ private:
#if defined(ENABLE_NPKIT) && (defined(ENABLE_NPKIT_EVENT_PRIM_LL_DATA_PROCESS_ENTRY) && defined(ENABLE_NPKIT_EVENT_PRIM_LL_DATA_PROCESS_EXIT) || defined(ENABLE_NPKIT_PRIM_COLLECT_DATA_PROCESS_TIME))
npkitWaitRecvSpins++;
#endif
if (checkAbort(spins, 0)) break;
if (checkAbort(abort, 1, spins)) break;
} while ((i4.flag1 != flag) || (i4.flag2 != flag));
uint64_t val64 = (uint64_t)(i4.data1) + (((uint64_t)i4.data2) << 32);
#else
@@ -177,7 +180,7 @@ private:
#if defined(ENABLE_NPKIT) && (defined(ENABLE_NPKIT_EVENT_PRIM_LL_DATA_PROCESS_ENTRY) && defined(ENABLE_NPKIT_EVENT_PRIM_LL_DATA_PROCESS_EXIT) || defined(ENABLE_NPKIT_PRIM_COLLECT_DATA_PROCESS_TIME))
npkitWaitRecvSpins++;
#endif
if (checkAbort(spins, 0)) break;
if (checkAbort(abort, 1, spins)) break;
} while ((flag1 != flag) || (flag2 != flag));
uint64_t val64 = data1 + (((uint64_t)data2) << 32);
#endif
@@ -241,7 +244,7 @@ private:
#if defined(ENABLE_NPKIT) && (defined(ENABLE_NPKIT_EVENT_PRIM_LL_DATA_PROCESS_ENTRY) && defined(ENABLE_NPKIT_EVENT_PRIM_LL_DATA_PROCESS_EXIT) || defined(ENABLE_NPKIT_PRIM_COLLECT_DATA_PROCESS_TIME))
npkitWaitRecvSpins++;
#endif
if (checkAbort(spins, 0)) break;
if (checkAbort(abort, 1, spins)) break;
} while(line[i].flag1 != flag || line[i].flag2 != flag);
uint64_t val64 = line[i].data1 + (((uint64_t)line[i].data2) << 32);
+12 -10
View File
@@ -86,16 +86,18 @@ private:
#endif
}
uint32_t abort = 0;
uint32_t* sync;
int abort = 0;
inline __device__ int checkAbort(int &spins, int i, int send) {
spins++;
if (abort == 0 && spins == NCCL_SPINS_BEFORE_CHECK_ABORT) {
abort = __atomic_load_n(ncclShmem.comm.abortFlag, __ATOMIC_SEQ_CST);
__device__ inline int checkAbort(int &abortCache, const int abortValue, int &spins) {
if (abortCache == 0 && ++spins == NCCL_SPINS_BEFORE_CHECK_ABORT) {
int abort = __atomic_load_n((ncclShmem.comm.abortFlag), __ATOMIC_SEQ_CST);
spins = 0;
if (abort) {
__atomic_store_n(&ncclShmem.aborted, abort, __ATOMIC_SEQ_CST);
abortCache |= abortValue;
}
}
return abort;
return abortCache;
}
inline __device__ void waitSend(int nbytes) {
@@ -104,7 +106,7 @@ private:
while (sendConnHeadCache + NCCL_STEPS < sendConnHead + 1) {
__builtin_amdgcn_s_sleep(1);
sendConnHeadCache = __atomic_load_n(sendConnHeadPtr, __ATOMIC_RELAXED);
if (checkAbort(spins, wid, 1)) break;
if (checkAbort(abort, 1, spins)) break;
}
if (sendConnFifo) {
sendConnFifo[sendStep[wid]%NCCL_STEPS].size = nbytes;
@@ -241,7 +243,7 @@ private:
load128(ptr+u*WARP_SIZE, vr[u], vr[u+1]);
needReload |= flagThread && (vr[u+1] != flag);
}
needReload &= (0 == checkAbort(spins, 0, 0));
needReload &= (0 == checkAbort(abort, 1, spins));
} while (__any(needReload));
#pragma unroll
for (int u=0; u<ELEMS_PER_THREAD; u+=2)
@@ -287,7 +289,7 @@ private:
load128(ptr+u*WARP_SIZE, vr[u], vr[u+1]);
needReload |= flagThread && (vr[u+1] != flag);
}
needReload &= (0 == checkAbort(spins, i, 0));
needReload &= (0 == checkAbort(abort, 1, spins));
} while (__any(needReload));
#pragma unroll
+243 -160
View File
@@ -59,7 +59,7 @@ class Primitives<
uint64_t connStepCache; // Cache last seen value of (*connStepPtr)
int connStepSize; // Connection step size
void* netDeviceHandle;
uint64_t accSize; // Accumulated size. Used by PAT operations
uint64_t accSize;
uint32_t* next_hdp_reg;
uint64_t* barriers;
uint64_t barrier_next = 0;
@@ -86,19 +86,21 @@ private:
#endif
}
inline __device__ void subBarrier() {
if (nworkers == WARP_SIZE) __syncwarp();
else
barrier();
}
inline __device__ void patBarrier() {
barrier();
}
inline __device__ bool checkAbort(int &spins) {
spins++;
if (!(flags & Aborted) && spins == NCCL_SPINS_BEFORE_CHECK_ABORT) {
if (__atomic_load_n(ncclShmem.comm.abortFlag, __ATOMIC_SEQ_CST)) {
flags |= Aborted;
ncclShmem.aborted = 1;
}
spins = 0;
}
return flags & Aborted;
inline __device__ void barrierAny() {
barrier();
}
inline __device__ void subBarrierAny() {
barrier();
}
inline __device__ uint64_t loadStepValue(uint64_t* ptr) {
@@ -129,7 +131,7 @@ private:
while (connStepCache + (isSendNotRecv ? NCCL_STEPS : 0) < step + StepPerSlice) {
__builtin_amdgcn_s_sleep(1);
connStepCache = loadStepValue(connStepPtr);
if (checkAbort(spins)) break;
if (checkAbort(flags, Aborted, spins)) break;
//if (spins == 0) printf("r=%d b=%d t=%d SPUN OUT got=%d want=%d\n", ncclShmem.comm.rank, blockIdx.x, threadIdx.x, int(connStepCache + (isSendNotRecv ? NCCL_STEPS : 0)), int(step+StepPerSlice));
if (spins == 0 && repeat > 0) {
repeat --;
@@ -482,13 +484,8 @@ public:
peerPtr->recv[connIndex].step += steps;
st_relaxed_sys_global(peerPtr->recv[connIndex].head, peerPtr->recv[connIndex].step);
while (ld_volatile_global(peerPtr->recv[connIndex].tail) < peerPtr->recv[connIndex].step) {
if (spins++ == NCCL_SPINS_BEFORE_CHECK_ABORT) {
if (*ncclShmem.comm.abortFlag) {
ncclShmem.aborted = 1;
break;
}
spins = 0;
}
int abort = 0;
if (checkAbort(abort, 1, spins)) break;
}
}
@@ -503,7 +500,7 @@ public:
int spins = 0;
while (connStepCache + (isSendNotRecv ? NCCL_STEPS : 0) < step + StepPerSlice) {
connStepCache = loadStepValue(connStepPtr);
if (checkAbort(spins)) break;
if (checkAbort(flags, Aborted, spins)) break;
}
void **ptrs = isSendNotRecv ? ncclShmem.groups[group].dsts
: ncclShmem.groups[group].srcs;
@@ -754,6 +751,9 @@ public:
flags = 0;
index = -1;
if (mode == primsModeDefault) { // Connect to ranks in sendPeers/recvPeers
// // For send operations, we need an extra warp to overlap the threadfence and the copy
// this->nworkers = nthreads - (MaxSend > 0 && nthreads >= NCCL_SIMPLE_EXTRA_GROUP_IF_NTHREADS_GE ? WARP_SIZE : 0);
int nrecv=0, nsend=0;
// Yes, for some template arguments this code will be unreachable. That's fine.
// coverity[dead_error_line]
@@ -783,68 +783,84 @@ public:
if (flags & (RoleWaitRecv|RolePostRecv)) peer = recvPeers[index];
if (flags & (RoleWaitSend|RolePostSend)) peer = sendPeers[index];
// Coverity thinks that index could be -1 here but that's not actually the case.
// coverity[negative_returns:FALSE]
int sendIpcReg;
int recvIpcReg;
int sendNetReg;
int recvNetReg;
if (P2p) {
sendIpcReg = p2pWork ? p2pWork->sendIpcReg : 0;
recvIpcReg = p2pWork ? p2pWork->recvIpcReg : 0;
sendNetReg = p2pWork ? p2pWork->sendNetReg : 0;
recvNetReg = p2pWork ? p2pWork->recvNetReg : 0;
} else {
recvIpcReg = sendIpcReg = collWork ? collWork->regUsed : 0;
recvNetReg = sendNetReg = collWork ? collWork->netRegUsed : 0;
}
// coverity[overrun-call] => Coverity think prims.index can be greater than 1
if (flags & (RoleWaitRecv|RolePostRecv)) loadRecvConn(ncclShmem.channel.peers[peer], connIndexRecv, collWork ? collWork->direct : 0, recvIpcReg, recvNetReg);
// coverity[overrun-call] => Coverity think prims.index can be greater than 1
if (flags & (RoleWaitSend|RolePostSend)) loadSendConn(ncclShmem.channel.peers[peer], connIndexSend, collWork ? collWork->direct : 0, sendIpcReg, sendNetReg);
// if (barrierAny(flags & NetDeviceUnpack)) {
// flags |= AnyNetDeviceUnpack;
// // RoleWaitRecv starts at tid=0, so this creates the bitmask of which recv peers
// // have NetDeviceUnpack.
// uint32_t mask = __ballot_sync(~0u, ((flags & RoleWaitRecv) && (flags & NetDeviceUnpack)) ? 1 : 0);
// if (tid == 0) {
// ncclShmem.groups[this->group].devicePlugin.unpack.unpackNetDeviceIndexMask = mask;
// }
// }
// coverity[negative_returns:FALSE] => coverity thinks that index could be -1 but that's not actually the case
// coverity[var_deref_model] => coverity thinks work can dereferenced if NULL but this is not the case
setDataPtrs(inputBuf, outputBuf, redOpArg, (struct ncclDevWorkCollReg*)collWork, sendIpcReg || recvIpcReg, peer);
// coverity[uninit_member] => coverity thinks fan.n is not initialized
} else if (mode == primsModePatRs || mode == primsModePatAg) { // Connect to all ranks +/- 2^n
flags |= PatMode;
accSize = 0;
const int roles[5] = { RoleWaitRecv, RolePostRecv, RoleWaitSend, RolePostSend, RoleInput | RoleOutput };
if (tid < 5) flags |= roles[tid];
int nranks = ncclShmem.comm.nRanks;
int rank = ncclShmem.comm.rank;
// A thread is responsible for rank +/- 2 ^ (tid%32). That should be fine as long as rank is a 32-bits integer.
index = tid % 32;
uint32_t delta = 1 << index;
const int roles[4] = { RoleWaitRecv, RoleWaitSend, RolePostSend, RolePostRecv};
int block = tid / 32;
if (block < 4 && delta < nranks) {
int role = roles[block];
if (mode == primsModePatRs) {
if (role & (RoleWaitRecv|RolePostRecv)) peer = (rank - delta + nranks) % nranks;
if (role & (RoleWaitSend|RolePostSend)) peer = (rank + delta) % nranks;
} else if (mode == primsModePatAg) {
if (role & (RoleWaitSend|RolePostSend)) peer = (rank - delta + nranks) % nranks;
if (role & (RoleWaitRecv|RolePostRecv)) peer = (rank + delta) % nranks;
}
flags |= role;
} else if (tid == 128) {
flags |= RoleInput | RoleOutput; // Only one will be used depending on the operation
if (tid < 32 && ((1UL<<tid) < nranks)) {
int rank = ncclShmem.comm.rank;
uint32_t delta = 1 << tid;
// Load recv peer
int recvPeer = mode == primsModePatRs ? (rank - delta + nranks) % nranks : (rank + delta) % nranks;
struct ncclPatPeer* peer = ((struct ncclPatPeer*)recvPeers)+tid;
struct ncclConnInfo* conn = peer->conn = ncclShmem.channel.peers[recvPeer]->recv+connIndexRecv;
peer->step = conn->step;
peer->buff = conn->buffs[NCCL_PROTO_SIMPLE];
peer->stepCache = loadStepValue(peer->tailPtr = conn->tail);
peer->headPtr = conn->head;
peer->accSize = 0;
peer->connStepSize = conn->stepSize/sizeof(T);
// Load send peer
int sendPeer = mode == primsModePatAg ? (rank - delta + nranks) % nranks : (rank + delta) % nranks;
peer = ((struct ncclPatPeer*)sendPeers)+tid;
conn = peer->conn = ncclShmem.channel.peers[sendPeer]->send+connIndexSend;
peer->step = conn->step;
peer->connFifo = conn->connFifo;
peer->buff = conn->buffs[NCCL_PROTO_SIMPLE];
peer->stepCache = loadStepValue(peer->headPtr = conn->head);
peer->tailPtr = conn->tail;
peer->accSize = 0;
peer->connStepSize = conn->stepSize/sizeof(T);
}
if (tid==0) {
ncclShmem.groups[group].userInput = (void*)inputBuf;
ncclShmem.groups[group].userOutput = (void*)outputBuf;
ncclShmem.redOpArgs[0] = redOpArg; // scaler for local input
}
patBarrier();
}
// Coverity thinks that index could be -1 here but that's not actually the case.
// coverity[negative_returns:FALSE]
int sendIpcReg;
int recvIpcReg;
int sendNetReg;
int recvNetReg;
if (P2p) {
sendIpcReg = p2pWork ? p2pWork->sendIpcReg : 0;
recvIpcReg = p2pWork ? p2pWork->recvIpcReg : 0;
sendNetReg = p2pWork ? p2pWork->sendNetReg : 0;
recvNetReg = p2pWork ? p2pWork->recvNetReg : 0;
} else {
recvIpcReg = sendIpcReg = collWork ? collWork->regUsed : 0;
recvNetReg = sendNetReg = collWork ? collWork->netRegUsed : 0;
}
// coverity[overrun-call] => Coverity think prims.index can be greater than 1
if (flags & (RoleWaitRecv|RolePostRecv)) loadRecvConn(ncclShmem.channel.peers[peer], connIndexRecv, collWork ? collWork->direct : 0, recvIpcReg, recvNetReg);
// coverity[overrun-call] => Coverity think prims.index can be greater than 1
if (flags & (RoleWaitSend|RolePostSend)) loadSendConn(ncclShmem.channel.peers[peer], connIndexSend, collWork ? collWork->direct : 0, sendIpcReg, sendNetReg);
// if (barrierAny(flags & NetDeviceUnpack)) {
// flags |= AnyNetDeviceUnpack;
// // RoleWaitRecv starts at tid=0, so this creates the bitmask of which recv peers
// // have NetDeviceUnpack.
// uint32_t mask = __ballot_sync(~0u, ((flags & RoleWaitRecv) && (flags & NetDeviceUnpack)) ? 1 : 0);
// if (tid == 0) {
// ncclShmem.groups[this->group].devicePlugin.unpack.unpackNetDeviceIndexMask = mask;
// }
// }
// coverity[negative_returns:FALSE] => coverity thinks that index could be -1 but that's not actually the case
// coverity[var_deref_model] => coverity thinks work can dereferenced if NULL but this is not the case
setDataPtrs(inputBuf, outputBuf, redOpArg, (struct ncclDevWorkCollReg*)collWork, sendIpcReg || recvIpcReg, peer);
// coverity[uninit_member] => coverity thinks fan.n is not initialized
}
__forceinline__ __device__ ~Primitives() {
if (flags&PatMode) return;
// Save steps for the next operation
if (flags & (RolePostSend|RolePostRecv)) conn->step = step;
if ((flags & NetRegMode) && (flags & RoleWaitSend)) {
@@ -854,7 +870,7 @@ public:
uint64_t prevStep = step - StepPerSlice;
volatile ssize_t* ptr = &(connFifo[prevStep%NCCL_STEPS].size);
int spins = 0;
while (*ptr != -1) if (checkAbort(spins)) break;
while (*ptr != -1) if (checkAbort(flags, Aborted, spins)) break;
}
if (flags & NetDeviceUnpack) {
@@ -872,7 +888,7 @@ public:
int spins = 0;
volatile uint64_t* tail = conn->tail;
volatile uint64_t* head = conn->head;
while (*tail > *head) if (checkAbort(spins)) break;
while (*tail > *head) if (checkAbort(flags, Aborted, spins)) break;
}
}
@@ -895,7 +911,7 @@ public:
if (slot) {
T* exchgPtr;
directBuff = (T*)outputBuf;
while ((void *)atomicAdd((unsigned long long *) slot,0) != nullptr && !checkAbort(spins));
while ((void *)atomicAdd((unsigned long long *) slot,0) != nullptr && !checkAbort(flags, Aborted, spins));
if (P2p) {
exchgPtr = (T*)outputBuf;
} else {
@@ -912,7 +928,7 @@ public:
void* ptr;
while (slot) {
ptr = (void *)atomicAdd((unsigned long long *) slot,0);
if (ptr != nullptr || checkAbort(spins)) break;
if (ptr != nullptr || checkAbort(flags, Aborted, spins)) break;
}
if (slot) {
@@ -931,7 +947,7 @@ public:
// Wait for consumer to consume previous value before trampling it.
if (slot && argSlot0 && argSlot1) {
T* exchgPtr;
while (((void *)atomicAdd((unsigned long long *) slot,0) != nullptr || *argSlot0 != 0 || *argSlot1 != 0) && !checkAbort(spins));
while (((void *)atomicAdd((unsigned long long *) slot,0) != nullptr || *argSlot0 != 0 || *argSlot1 != 0) && !checkAbort(flags, Aborted, spins));
// If there is no recv, then we are directly pulling from input buffer (e.g. directScatter)
// Otherwise, we are pulling from output buffer (e.g. recvCopyDirectSend)
directBuff = MaxRecv == 0 ? (T*)inputBuf : (T*)outputBuf;
@@ -961,7 +977,7 @@ public:
void* ptr;
while (slot) {
ptr = (void *)atomicAdd((unsigned long long *) slot,0);
if (ptr != nullptr || checkAbort(spins)) break;
if (ptr != nullptr || checkAbort(flags, Aborted, spins)) break;
}
if (slot && argSlot0 && argSlot1) {
@@ -972,7 +988,7 @@ public:
while (true) {
arg0 = *argSlot0;
arg1 = *argSlot1;
if ((arg0 != 0 && arg1 != 0) || checkAbort(spins)) break;
if ((arg0 != 0 && arg1 != 0) || checkAbort(flags, Aborted, spins)) break;
}
ncclShmem.redOpArgs[1 + index] = ((arg1 & 0xffffffff) << 32) | (arg0 & 0xffffffff);
}
@@ -1020,8 +1036,8 @@ public:
__device__ __forceinline__ void recv(intptr_t outIx, int eltN, bool postOp=false) {
genericOp<0, 0, 1, 0, -1, Output>(-1, outIx, eltN, postOp);
}
__device__ __forceinline__ void directRecv(intptr_t inpIx, intptr_t outIx, int eltN, bool postOp=false) {
genericOp<1, 0, 1, 0, -1, Output>(inpIx, outIx, eltN, postOp);
__device__ __forceinline__ void directRecv(intptr_t outIx, int eltN, bool postOp=false) {
genericOp<1, 0, 1, 0, -1, Output>(outIx, outIx, eltN, postOp);
}
__device__ __forceinline__ void directRecvCopy(intptr_t inpIx, intptr_t outIx, int eltN) {
genericOp<1, 0, 1, 0, -1, Output>(inpIx, outIx, eltN, /*postOp=*/false);
@@ -1099,54 +1115,65 @@ public:
ScatterGatherOp<1, 0, 1, 0>(-1, outIx, totalElem, peerElem, peerOffset, skip, shift, /*postOp=*/false);
}
__device__ __forceinline__ void patReduce(int recvPow2, int sendPow2, intptr_t inpIx, intptr_t outIx, int recvOffset, int sendOffset, int sendStepOffset, int nelem, int postRecv, int postSend) {
nelem = nelem < 0 ? 0 : nelem;
__device__ __forceinline__ void patReduce(struct ncclPatStep* ps, struct ncclPatShmem* shmem) {
if (ps->flags & PatSkipped) { patBarrier(); patBarrier(); return; } // Skipped
int nelem = ps->nelem < 0 ? 0 : ps->nelem;
T* userInput = (T*)ncclShmem.groups[group].userInput;
T* userOutput = (T*)ncclShmem.groups[group].userOutput;
if (recvPow2 >= 0 && recvPow2 == index && (flags & RoleWaitRecv)) {
ncclShmem.groups[group].srcs[0] = (T*)(connEltsFifo + (step%NCCL_STEPS)*connStepSize) + recvOffset;
int spins = 0;
while (connStepCache < step + StepPerSlice) {
connStepCache = loadStepValue(connStepPtr);
if (checkAbort(spins)) break;
}
if (postRecv) step += StepPerSlice;
bool recv = ps->recvDim >= 0 && (flags & (RolePostRecv|RoleWaitRecv));
bool send = ps->sendDim >= 0 && (flags & (RolePostSend|RoleWaitSend));
bool postRecv = ps->postRecv && recv;
bool postSend = ps->postSend && send;
struct ncclPatPeer* peer = NULL;
if (recv) {
peer = shmem->recvDims+ps->recvDim;
step = peer->step;
}
if (sendPow2 >= 0 && sendPow2 == index && (flags & RoleWaitSend)) {
int spins = 0;
while (connStepCache + NCCL_STEPS < step + sendStepOffset + StepPerSlice) {
connStepCache = loadStepValue(connStepPtr);
if (checkAbort(spins)) break;
}
ncclShmem.groups[group].dsts[0] = (T*)(connEltsFifo + ((step+sendStepOffset)%NCCL_STEPS)*connStepSize) + sendOffset;
if (accSize < sendOffset + nelem + (step+sendStepOffset)*connStepSize) {
// New data, add our own data to it.
ncclShmem.groups[group].srcs[1] = userInput + inpIx;
accSize = sendOffset + nelem + (step+sendStepOffset)*connStepSize;
if (flags & ConnFifoEnabled)
connFifo[(step+sendStepOffset)%NCCL_STEPS].size = (sendOffset + nelem)*sizeof(T);
} else {
// There is already data in there, accumulate instead of writing to it.
ncclShmem.groups[group].srcs[1] = ncclShmem.groups[group].dsts[0];
}
if (postSend) step += StepPerSlice;
if (send) {
peer = shmem->sendDims+ps->sendDim;
step = peer->step;
}
if (sendPow2 < 0 && (flags & RoleOutput)) { // Destination is our own local buffer
ncclShmem.groups[group].dsts[0] = userOutput + outIx;
if (accSize < outIx + nelem) {
if (recv && (flags & RoleWaitRecv)) {
ncclShmem.groups[group].srcs[0] = ((T*)peer->buff) + (step%NCCL_STEPS)*peer->connStepSize + ps->recvOffset;
int spins = 0;
while (peer->stepCache < step + StepPerSlice) {
peer->stepCache = loadStepValue(peer->tailPtr);
if (checkAbort(flags, Aborted, spins)) break;
}
}
if (send && (flags & RoleWaitSend)) {
int spins = 0;
while (peer->stepCache + NCCL_STEPS < step + ps->stepOffset + StepPerSlice) {
peer->stepCache = loadStepValue(peer->headPtr);
if (checkAbort(flags, Aborted, spins)) break;
}
ncclShmem.groups[group].dsts[0] = ((T*)peer->buff) + ((step+ps->stepOffset)%NCCL_STEPS)*peer->connStepSize + ps->sendOffset;
if (peer->accSize < ps->sendOffset + nelem + (step+ps->stepOffset)*peer->connStepSize) {
// New data, add our own data to it.
ncclShmem.groups[group].srcs[1] = userInput + inpIx;
accSize = outIx + nelem;
ncclShmem.groups[group].srcs[1] = userInput + ps->inpIx;
} else {
// There is already data in there, accumulate instead of writing to it.
ncclShmem.groups[group].srcs[1] = ncclShmem.groups[group].dsts[0];
}
}
barrier();
long long int localAccSize = shmem->localAccSize;
if (ps->sendDim < 0 && (flags & RoleOutput)) { // Destination is our own local buffer
ncclShmem.groups[group].dsts[0] = userOutput + ps->outIx;
if (localAccSize < ps->outIx + nelem) {
// New data, add our own data to it.
ncclShmem.groups[group].srcs[1] = userInput + ps->inpIx;
localAccSize = ps->outIx + nelem;
} else {
// There is already data in there, accumulate instead of writing to it.
ncclShmem.groups[group].srcs[1] = ncclShmem.groups[group].dsts[0];
}
}
patBarrier();
int nSrcs = 2;
void** srcs = ncclShmem.groups[group].srcs;
if (recvPow2 < 0) { srcs++; nSrcs--; } // No peer to receive from, remove one source
if (ps->recvDim < 0) { srcs++; nSrcs--; } // No peer to receive from, remove one source
int workSize = ncclShmem.aborted ? 0 : nelem;
@@ -1154,59 +1181,92 @@ public:
(tid, nthreads, ncclShmem.redOpArgs[0], nullptr, /*postOp=*/false,
nSrcs, srcs, 1, ncclShmem.groups[group].dsts, workSize);
barrier();
if (postRecv && recvPow2 >= 0 && recvPow2 == index && (flags & RolePostRecv)) postPeer<1, 0>(0 < nelem);
if (postSend && sendPow2 >= 0 && sendPow2 == index && (flags & RolePostSend)) postPeer<0, 1>(0 < nelem);
// Store conn step here inside the two barriers to make sure next reload will see the update.
if (postSend && (flags & RolePostSend)) {
if (peer->connFifo) {
peer->connFifo[step%NCCL_STEPS].size = (ps->sendOffset + nelem)*sizeof(T);
}
peer->step = step += StepPerSlice;
st_relaxed_sys_global(&peer->conn->step, step);
}
if (postRecv && (flags & RolePostRecv)) {
peer->step = step += StepPerSlice;
st_relaxed_sys_global(&peer->conn->step, step); // Also save in global mem for next op
}
// Update accSize
if (ps->sendDim < 0 && (flags & RoleOutput)) atomicMax(&shmem->localAccSize, localAccSize);
if (ps->sendDim >= 0 && (flags & RoleWaitSend)) atomicMax(&peer->accSize, ps->sendOffset + nelem + (step+ps->stepOffset)*peer->connStepSize);
patBarrier();
if (postSend && (flags & RolePostSend)) {
if (nelem > 0 || peer->connFifo) fence_acq_rel_sys();
st_relaxed_sys_global(peer->tailPtr, step);
}
if (postRecv && (flags & RolePostRecv)) {
st_relaxed_sys_global(peer->headPtr, step);
}
}
__device__ __forceinline__ void patCopy(int recvPow2, int sendPow2, intptr_t inpIx, intptr_t outIx, int recvOffset, int sendOffset, int recvStepOffset, int nelem, int postRecv, int postSend) {
nelem = nelem < 0 ? 0 : nelem;
__device__ __forceinline__ void patCopy(struct ncclPatStep* ps, struct ncclPatShmem* shmem) {
if (ps->flags & PatSkipped) { patBarrier(); patBarrier(); return; } // Skipped
int nelem = ps->nelem < 0 ? 0 : ps->nelem;
T* userInput = (T*)ncclShmem.groups[group].userInput;
T* userOutput = (T*)ncclShmem.groups[group].userOutput;
if (recvPow2 >= 0 && recvPow2 == index && (flags & RoleWaitRecv)) {
ncclShmem.groups[group].srcs[0] = (T*)(connEltsFifo + ((step+recvStepOffset)%NCCL_STEPS)*connStepSize) + recvOffset;
int spins = 0;
while (connStepCache < step + recvStepOffset + StepPerSlice) {
connStepCache = loadStepValue(connStepPtr);
if (checkAbort(spins)) break;
}
if (accSize < recvOffset + nelem + (step+recvStepOffset)*connStepSize) {
// New data, copy to our output buffer.
ncclShmem.groups[group].dsts[1] = userOutput + outIx;
accSize = recvOffset + nelem + (step+recvStepOffset)*connStepSize;
} else {
ncclShmem.groups[group].dsts[1] = ncclShmem.groups[group].srcs[0]; // Already done
}
if (postRecv) step += StepPerSlice;
bool recv = ps->recvDim >= 0 && (flags & (RolePostRecv|RoleWaitRecv));
bool send = ps->sendDim >= 0 && (flags & (RolePostSend|RoleWaitSend));
bool postRecv = ps->postRecv && recv;
bool postSend = ps->postSend && send;
struct ncclPatPeer* peer = NULL;
if (recv) {
peer = shmem->recvDims+ps->recvDim;
step = peer->step;
}
if (sendPow2 >= 0 && sendPow2 == index && (flags & RoleWaitSend)) {
int spins = 0;
while (connStepCache + NCCL_STEPS < step + StepPerSlice) {
connStepCache = loadStepValue(connStepPtr);
if (checkAbort(spins)) break;
}
ncclShmem.groups[group].dsts[0] = (T*)(connEltsFifo + (step%NCCL_STEPS)*connStepSize) + sendOffset;
if (postSend) {
if (flags & ConnFifoEnabled)
connFifo[step%NCCL_STEPS].size = (sendOffset + nelem)*sizeof(T);
step += StepPerSlice;
}
if (send) {
peer = shmem->sendDims+ps->sendDim;
step = peer->step;
}
if (recvPow2 < 0 && (flags & RoleInput)) { // Source is our own local buffer
ncclShmem.groups[group].srcs[0] = userInput + inpIx;
if (accSize < inpIx + nelem) {
if (recv && (flags & RoleWaitRecv)) {
ncclShmem.groups[group].srcs[0] = ((T*)peer->buff) + ((step+ps->stepOffset)%NCCL_STEPS)*peer->connStepSize + ps->recvOffset;
int spins = 0;
while (peer->stepCache < step + ps->stepOffset + StepPerSlice) {
peer->stepCache = loadStepValue(peer->tailPtr);
if (checkAbort(flags, Aborted, spins)) break;
}
if (peer->accSize < ps->recvOffset + nelem + (step+ps->stepOffset)*peer->connStepSize) {
// New data, copy to our output buffer.
ncclShmem.groups[group].dsts[1] = userOutput + outIx;
accSize = inpIx + nelem;
ncclShmem.groups[group].dsts[1] = userOutput + ps->outIx;
} else {
ncclShmem.groups[group].dsts[1] = ncclShmem.groups[group].srcs[0]; // Already done
}
}
barrier();
if (send && (flags & RoleWaitSend)) {
int spins = 0;
while (peer->stepCache + NCCL_STEPS < step + StepPerSlice) {
peer->stepCache = loadStepValue(peer->headPtr);
if (checkAbort(flags, Aborted, spins)) break;
}
ncclShmem.groups[group].dsts[0] = ((T*)peer->buff) + (step%NCCL_STEPS)*peer->connStepSize + ps->sendOffset;
}
long long int localAccSize = shmem->localAccSize;
if (ps->recvDim < 0 && (flags & RoleInput)) { // Source is our own local buffer
ncclShmem.groups[group].srcs[0] = userInput + ps->inpIx;
if (localAccSize < ps->inpIx + nelem) {
// New data, copy to our output buffer.
ncclShmem.groups[group].dsts[1] = userOutput + ps->outIx;
localAccSize = ps->inpIx + nelem;
} else {
// Already done
ncclShmem.groups[group].dsts[1] = ncclShmem.groups[group].srcs[0];
}
}
patBarrier();
int nDsts = 2;
void** dsts = ncclShmem.groups[group].dsts;
if (sendPow2 < 0) { dsts++; nDsts--; } // No peer to send to, remove one dest
if (ps->sendDim < 0) { dsts++; nDsts--; } // No peer to send to, remove one dest
if (ncclShmem.groups[group].srcs[0] == ncclShmem.groups[group].dsts[1]) nDsts--; // In-place or already done.
int workSize = ncclShmem.aborted ? 0 : nelem;
@@ -1215,9 +1275,32 @@ public:
(tid, nthreads, ncclShmem.redOpArgs[0], nullptr, /*postOp=*/false,
1, ncclShmem.groups[group].srcs, nDsts, dsts, workSize);
barrier();
if (postRecv && recvPow2 >= 0 && recvPow2 == index && (flags & RolePostRecv)) postPeer<1, 0>(0 < nelem);
if (postSend && sendPow2 >= 0 && sendPow2 == index && (flags & RolePostSend)) postPeer<0, 1>(0 < nelem);
// Store conn step here inside the two barriers to make sure next reload will see the update.
if (postSend && (flags & RolePostSend)) {
if (peer->connFifo) {
peer->connFifo[step%NCCL_STEPS].size = (ps->sendOffset + nelem)*sizeof(T);
}
peer->step = step += StepPerSlice;
st_relaxed_sys_global(&peer->conn->step, step);
}
if (postRecv && (flags & RolePostRecv)) {
peer->step = step += StepPerSlice;
st_relaxed_sys_global(&peer->conn->step, step); // Also save in global mem for next op
}
// Update accSize
if (ps->recvDim < 0 && (flags & RoleInput)) atomicMax(&shmem->localAccSize, localAccSize);
if (ps->recvDim >= 0 && (flags & RoleWaitRecv)) atomicMax(&peer->accSize, ps->recvOffset + nelem + (step+ps->stepOffset)*peer->connStepSize);
patBarrier();
if (postSend && (flags & RolePostSend)) {
if (nelem > 0 || peer->connFifo) fence_acq_rel_sys();
st_relaxed_sys_global(peer->tailPtr, step);
}
if (postRecv && (flags & RolePostRecv)) {
st_relaxed_sys_global(peer->headPtr, step);
}
}
// MSCCL primitives
+49 -12
View File
@@ -170,29 +170,66 @@ struct RunWorkColl<ncclFuncReduceScatter, T, RedOp, NCCL_ALGO_RING, NCCL_PROTO_L
template<typename T, typename RedOp>
struct RunWorkColl<ncclFuncReduceScatter, T, RedOp, NCCL_ALGO_PAT, NCCL_PROTO_SIMPLE> {
__device__ __forceinline__ void run(int tid, int nthreads, struct ncclDevWorkColl* work) {
#if __CUDA_ARCH__ >= 600
using Proto = ProtoSimple<1, 1>;
const int nranks = ncclShmem.comm.nRanks;
const int rank = ncclShmem.comm.rank;
size_t count, channelOffset, channelCount, chunkCount;
ncclCollCbdPart(work, ncclShmem.channelId, Proto::Id, sizeof(T), &count, &channelOffset, &channelCount, &chunkCount);
T *inputBuf = (T*)work->sendbuff;
T *outputBuf = (T*)work->recvbuff;
Primitives<T, RedOp, FanSymmetric<1>, 0, Proto, 0> prims
(tid, nthreads, NULL, NULL, inputBuf, outputBuf, work->redOpArg, 0*Proto::MaxGroupWidth, 0, 0, nullptr, nullptr, 0, primsModePatRs);
static constexpr int nworkers = NCCL_PAT_NWORKERS;
struct ncclPatShmem* shmem = (struct ncclPatShmem*)ncclScratchForWarp(0);
uint64_t pollCount = 0;
__syncthreads(); // Don't start using shared mem until everyone arrives
for (int i=tid; i<NCCL_SHMEM_PAT_STEPS; i+=nthreads) shmem->patSteps[i].flags = 0;
if (tid == 0) shmem->localAccSize = 0;
if (tid == nworkers) shmem->parallelFactor = 0;
__syncthreads();
PatRSAlgorithm<T> patAlgo(chunkCount*sizeof(T), NCCL_STEPS, channelOffset, channelOffset + channelCount, count, chunkCount, rank, nranks);
int last = 0;
while (!last) {
int recvDim, sendDim, recvOffset, sendOffset, sendStepOffset, postRecv, postSend, nelem;
size_t inpIx, outIx;
patAlgo.getNextOp(recvDim, sendDim, inpIx, outIx, recvOffset, sendOffset, sendStepOffset, nelem, postRecv, postSend, last);
prims.patReduce(recvDim, sendDim, inpIx, outIx, recvOffset, sendOffset, sendStepOffset, nelem, postRecv, postSend);
if (tid == nworkers) { // Algo computation thread
PatRSAlgorithm<T> patAlgo(chunkCount*sizeof(T), NCCL_STEPS, NCCL_PAT_NWORKERS/WARP_SIZE, channelOffset, channelOffset + channelCount, count, chunkCount, rank, nranks);
int parallelFactor = shmem->parallelFactor = patAlgo.getParallelFactor();
int step = 0;
while (1) {
struct ncclPatStep* ps = shmem->patSteps+(step%NCCL_SHMEM_PAT_STEPS);
cuda::atomic_ref<int, cuda::thread_scope_block> poll(ps->flags);
while (poll.load(cuda::memory_order_acquire) != 0) pollCount++; // Wait for workers to be done with step 'step-NCCL_SHMEM_PAT_STEPS'
patAlgo.getNextOp(ps);
int last = ps->last;
step++;
if (last == 2) break;
}
} else if (tid < nworkers) { // Worker threads
T *inputBuf = (T*)work->sendbuff;
T *outputBuf = (T*)work->recvbuff;
int parallelFactor = 0;
volatile int* pfPtr = &shmem->parallelFactor;
while (parallelFactor == 0) parallelFactor = *pfPtr;
int groupSize = nworkers/(WARP_SIZE*parallelFactor) * WARP_SIZE;
int group = tid / groupSize;
int nGroups = nworkers / groupSize;
int tidInGroup = tid - group*groupSize;
// We don't use recvPeers/sendPeers so let's pass shmem structs instead
Primitives<T, RedOp, FanSymmetric<1>, 0, Proto, 0> prims
(tidInGroup, groupSize, (int*)shmem->recvDims, (int*)shmem->sendDims, inputBuf, outputBuf, work->redOpArg, group, 0, 0, nullptr, nullptr, 0, primsModePatRs);
int step = group;
while(1) {
struct ncclPatStep* ps = shmem->patSteps+(step%NCCL_SHMEM_PAT_STEPS);
cuda::atomic_ref<int, cuda::thread_scope_block> poll(ps->flags);
while (poll.load(cuda::memory_order_acquire) == 0) pollCount++; // Wait for compute thread
int last = ps->last;
prims.patReduce(ps, shmem);
if (tidInGroup == 0) poll.store(0, cuda::memory_order_release); // Return element to compute thread
if (last) break;
step += nGroups;
}
}
#endif
}
};
template<typename T, typename RedOp>
struct RunWorkColl<ncclFuncReduceScatter, T, RedOp, NCCL_ALGO_NVLS, NCCL_PROTO_SIMPLE> {
__device__ __forceinline__ void run(int tid, int/*nthreads*/, struct ncclDevWorkColl* work) {
+1 -1
View File
@@ -122,7 +122,7 @@ struct RunWorkBatch<ncclFuncSendRecv, T, RedOp, NCCL_ALGO_RING, NCCL_PROTO_SIMPL
size_t cursor = 0;
do {
int n = min(size_t(chunkSize), bytes-cursor);
prims.directRecv(cursor, cursor, n);
prims.directRecv(cursor, n);
cursor += n;
} while (cursor < bytes);