Merge remote-tracking branch 'nccl/master' into develop
This commit is contained in:
+50
-12
@@ -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
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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) {
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
Reference in New Issue
Block a user