2019-11-21 13:41:10 -08:00
|
|
|
/*************************************************************************
|
2020-05-12 14:40:18 -07:00
|
|
|
* Copyright (c) 2016-2020, NVIDIA CORPORATION. All rights reserved.
|
2020-01-15 17:54:27 -07:00
|
|
|
* Modifications Copyright (c) 2019-2020 Advanced Micro Devices, Inc. All rights reserved.
|
2019-11-21 13:41:10 -08:00
|
|
|
*
|
|
|
|
|
* See LICENSE.txt for license information
|
|
|
|
|
************************************************************************/
|
|
|
|
|
|
2020-07-08 11:06:50 -07:00
|
|
|
template <typename T, class FUNC, int NRECV, int NSEND>
|
|
|
|
|
class ncclLLPrimitives {
|
|
|
|
|
private:
|
|
|
|
|
const int tid;
|
|
|
|
|
const int nthreads;
|
|
|
|
|
const int wid;
|
|
|
|
|
const int stepLines;
|
|
|
|
|
int nrecv = 0;
|
|
|
|
|
int nsend = 0;
|
2019-11-19 14:57:39 -08:00
|
|
|
struct ncclConnInfo* recvConn = NULL;
|
|
|
|
|
volatile uint64_t* recvConnHeadPtr = NULL;
|
|
|
|
|
uint64_t recvConnHead;
|
|
|
|
|
|
|
|
|
|
struct ncclConnInfo* sendConn = NULL;
|
|
|
|
|
volatile int* sendConnFifoPtr = NULL;
|
|
|
|
|
volatile uint64_t* sendConnHeadPtr = NULL;
|
|
|
|
|
uint64_t sendConnHead;
|
|
|
|
|
uint64_t sendConnHeadCache; // Cache last seen value
|
2020-07-08 11:06:50 -07:00
|
|
|
|
|
|
|
|
uint64_t recvStep[NRECV];
|
2019-11-19 14:57:39 -08:00
|
|
|
uint64_t sendStep[NSEND];
|
2020-07-08 11:06:50 -07:00
|
|
|
union ncclLLFifoLine* recvBuff[NRECV];
|
2019-11-19 14:57:39 -08:00
|
|
|
union ncclLLFifoLine* sendBuff[NSEND];
|
|
|
|
|
struct ncclDevComm* comm;
|
|
|
|
|
|
2020-07-08 11:06:50 -07:00
|
|
|
inline __device__ int recvOffset(int i) { return (recvStep[i]%NCCL_STEPS)*stepLines; }
|
|
|
|
|
inline __device__ int sendOffset(int i) { return (sendStep[i]%NCCL_STEPS)*stepLines; }
|
|
|
|
|
inline __device__ union ncclLLFifoLine* recvPtr(int i) { return recvBuff[i]+recvOffset(i); }
|
|
|
|
|
inline __device__ union ncclLLFifoLine* sendPtr(int i) { return sendBuff[i]+sendOffset(i); }
|
|
|
|
|
inline __device__ uint32_t recvFlag(int i) { return NCCL_LL_FLAG(recvStep[i]+1); }
|
|
|
|
|
inline __device__ uint32_t sendFlag(int i) { return NCCL_LL_FLAG(sendStep[i]+1); }
|
2019-11-19 14:57:39 -08:00
|
|
|
|
|
|
|
|
inline __device__ void barrier() {
|
2019-11-21 13:41:10 -08:00
|
|
|
#if defined(__HIP_PLATFORM_HCC__) || defined(__HCC__) || defined(__HIPCC__)
|
|
|
|
|
__syncthreads();
|
|
|
|
|
#else
|
2020-07-08 11:06:50 -07:00
|
|
|
asm volatile ("basync 1, %0;" :: "r"(nthreads));
|
2019-11-21 13:41:10 -08:00
|
|
|
#endif
|
2019-11-19 14:57:39 -08:00
|
|
|
}
|
|
|
|
|
|
|
|
|
|
uint32_t spins = 0;
|
|
|
|
|
uint32_t abort = 0;
|
|
|
|
|
|
|
|
|
|
inline __device__ int checkAbort(int i, int send) {
|
|
|
|
|
spins++;
|
|
|
|
|
if (abort == 0 && spins == SPINS_BEFORE_CHECK_ABORT) {
|
2019-11-21 13:41:10 -08:00
|
|
|
abort = LOAD(comm->abortFlag);
|
2019-11-19 14:57:39 -08:00
|
|
|
spins = 0;
|
|
|
|
|
}
|
|
|
|
|
return abort;
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
inline __device__ void waitSend(int nbytes) {
|
|
|
|
|
spins = 0;
|
2020-07-08 11:06:50 -07:00
|
|
|
if (sendConnHeadPtr) {
|
|
|
|
|
while (sendConnHeadCache + NCCL_STEPS < sendConnHead + 1) {
|
|
|
|
|
sendConnHeadCache = LOAD(sendConnHeadPtr);
|
2019-11-19 14:57:39 -08:00
|
|
|
if (checkAbort(wid, 1)) break;
|
|
|
|
|
}
|
2020-07-08 11:06:50 -07:00
|
|
|
if (sendConnFifoPtr) {
|
|
|
|
|
int size = ((sendConnHead & NCCL_LL_CLEAN_MASK) == NCCL_LL_CLEAN_MASK) ? stepLines*sizeof(union ncclLLFifoLine) : nbytes;
|
|
|
|
|
STORE(sendConnFifoPtr+sendConnHead%NCCL_STEPS, size);
|
2019-11-19 14:57:39 -08:00
|
|
|
}
|
2020-07-08 11:06:50 -07:00
|
|
|
sendConnHead += 1;
|
2019-11-19 14:57:39 -08:00
|
|
|
}
|
|
|
|
|
barrier();
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
inline __device__ void incRecv(int i) {
|
2020-07-08 11:06:50 -07:00
|
|
|
recvStep[i] += 1;
|
2019-11-19 14:57:39 -08:00
|
|
|
}
|
|
|
|
|
inline __device__ void postRecv() {
|
|
|
|
|
barrier();
|
2020-07-08 11:06:50 -07:00
|
|
|
if (recvConnHeadPtr) STORE(recvConnHeadPtr, recvConnHead += 1);
|
2019-11-19 14:57:39 -08:00
|
|
|
}
|
|
|
|
|
|
|
|
|
|
inline __device__ void incSend(int i, int offset) {
|
|
|
|
|
// LL Cleanup : write all flags in the slice to make sure we don't have
|
2020-07-08 11:06:50 -07:00
|
|
|
// data corruption when flag loops ove
|
|
|
|
|
if ((sendStep[i] & NCCL_LL_CLEAN_MASK) == NCCL_LL_CLEAN_MASK) {
|
2020-05-12 14:40:18 -07:00
|
|
|
for (int o = offset; o<stepLines; o+=nthreads) storeLL(sendPtr(i)+o, 0, sendFlag(i));
|
2019-11-19 14:57:39 -08:00
|
|
|
}
|
2020-07-08 11:06:50 -07:00
|
|
|
sendStep[i]++;
|
2019-11-19 14:57:39 -08:00
|
|
|
}
|
|
|
|
|
|
|
|
|
|
__device__ uint64_t readLL(int i, int offset) {
|
|
|
|
|
union ncclLLFifoLine* src = recvPtr(i) + offset;
|
|
|
|
|
uint32_t flag = recvFlag(i);
|
|
|
|
|
uint32_t data1, flag1, data2, flag2;
|
|
|
|
|
spins = 0;
|
2019-11-21 13:41:10 -08:00
|
|
|
#if defined(__HIP_PLATFORM_HCC__) || defined(__HCC__) || defined(__HIPCC__)
|
|
|
|
|
using Vec = uint32_t __attribute__((ext_vector_type(4)));
|
|
|
|
|
Vec i4;
|
|
|
|
|
do {
|
2020-09-25 13:46:26 -06:00
|
|
|
asm volatile ("flat_load_dwordx4 %0, %1, glc, slc\n"
|
|
|
|
|
"s_waitcnt vmcnt(0)\n" : "=v"(i4) : "v"(src));
|
2020-02-25 13:41:02 -08:00
|
|
|
if (checkAbort(i, 0)) break;
|
2019-11-21 13:41:10 -08:00
|
|
|
} while ((i4[1] != flag) || (i4[3] != flag));
|
|
|
|
|
uint64_t val64 = (uint64_t)(i4[0]) + (((uint64_t)i4[2]) << 32);
|
|
|
|
|
#else
|
2019-11-19 14:57:39 -08:00
|
|
|
do {
|
|
|
|
|
asm volatile("ld.volatile.global.v4.u32 {%0,%1,%2,%3}, [%4];" : "=r"(data1), "=r"(flag1), "=r"(data2), "=r"(flag2) : "l"(&src->i4));
|
|
|
|
|
if (checkAbort(i, 0)) break;
|
|
|
|
|
} while ((flag1 != flag) || (flag2 != flag));
|
|
|
|
|
uint64_t val64 = data1 + (((uint64_t)data2) << 32);
|
2019-11-21 13:41:10 -08:00
|
|
|
#endif
|
2019-11-19 14:57:39 -08:00
|
|
|
return val64;
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
__device__ void storeLL(union ncclLLFifoLine* dst, uint64_t val, uint32_t flag) {
|
2019-11-21 13:41:10 -08:00
|
|
|
#if defined(__HIP_PLATFORM_HCC__) || defined(__HCC__) || defined(__HIPCC__)
|
|
|
|
|
using Vec = uint32_t __attribute__((ext_vector_type(4)));
|
|
|
|
|
Vec i4;
|
|
|
|
|
i4[0] = val & 0xffffffff;
|
|
|
|
|
i4[1] = flag;
|
|
|
|
|
i4[2] = (val >> 32);
|
|
|
|
|
i4[3] = flag;
|
2020-09-25 13:46:26 -06:00
|
|
|
asm volatile ("flat_store_dwordx4 %0, %1, glc, slc\n"
|
|
|
|
|
"s_waitcnt vmcnt(0)\n" : : "v"(dst), "v"(i4));
|
2019-11-21 13:41:10 -08:00
|
|
|
#else
|
2019-11-19 14:57:39 -08:00
|
|
|
asm volatile("st.volatile.global.v4.u32 [%0], {%1,%2,%3,%4};" :: "l"(&dst->i4), "r"((uint32_t)val), "r"(flag), "r"((uint32_t)(val >> 32)), "r"(flag));
|
2019-11-21 13:41:10 -08:00
|
|
|
#endif
|
2019-11-19 14:57:39 -08:00
|
|
|
}
|
|
|
|
|
|
2020-07-08 11:06:50 -07:00
|
|
|
// Using memcpy handles misaligned pointer
|
2019-11-19 14:57:39 -08:00
|
|
|
__device__ uint64_t readAL(uint64_t* src) {
|
|
|
|
|
uint64_t val;
|
|
|
|
|
memcpy((char*)&val, (char*)src, sizeof(uint64_t));
|
|
|
|
|
return val;
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
__device__ void storeAL(uint64_t* dst, uint64_t val, uint32_t nbytes) {
|
|
|
|
|
memcpy((char*)dst, (char*)&val, nbytes);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
template <int RECV, int SEND, int SRC, int DST>
|
|
|
|
|
__device__ void LLGenericOp(const T* srcPtr, T* dstPtr, int nelem) {
|
|
|
|
|
uint32_t nbytes = nelem < 0 ? 0 : nelem*sizeof(T);
|
|
|
|
|
uint32_t npack = DIVUP(nbytes, sizeof(uint64_t));
|
|
|
|
|
uint64_t* srcPack = (uint64_t*)srcPtr;
|
|
|
|
|
uint64_t* dstPack = (uint64_t*)dstPtr;
|
|
|
|
|
int offset = tid;
|
|
|
|
|
|
|
|
|
|
// Always waitSend in case of cleanup
|
|
|
|
|
if (SEND) waitSend(npack*sizeof(union ncclLLFifoLine));
|
|
|
|
|
|
|
|
|
|
// Do multiples of 64 bits
|
2019-11-21 13:41:10 -08:00
|
|
|
#pragma unroll 1
|
2019-11-19 14:57:39 -08:00
|
|
|
for (; offset<npack; offset+=nthreads) {
|
|
|
|
|
// Recv : local, then intra-node, then inter-node
|
|
|
|
|
uint64_t val = SRC ? readAL(srcPack+offset) : readLL(0, offset);
|
|
|
|
|
if (RECV) {
|
|
|
|
|
if (SRC) val = MULTI<FUNC, T>()(readLL(0, offset), val);
|
|
|
|
|
for (int i=1; i<NRECV && i<nrecv; i++) {
|
|
|
|
|
val = MULTI<FUNC, T>()(readLL(i, offset), val);
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// Send : inter-node, then intra-node, then local
|
|
|
|
|
if (SEND) {
|
|
|
|
|
for (int i=1; i<NSEND && i<nsend; i++) storeLL(sendPtr(i)+offset, val, sendFlag(i));
|
|
|
|
|
storeLL(sendPtr(0)+offset, val, sendFlag(0));
|
|
|
|
|
}
|
|
|
|
|
if (DST) {
|
|
|
|
|
if (((offset*sizeof(uint64_t)) ^ nbytes) < sizeof(uint64_t)) {
|
|
|
|
|
// Last incomplete word
|
|
|
|
|
storeAL(dstPack+offset, val, nbytes & 0x7);
|
|
|
|
|
} else {
|
|
|
|
|
storeAL(dstPack+offset, val, sizeof(uint64_t));
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
FOR_RECV(incRecv); if (RECV) postRecv();
|
|
|
|
|
FOR_SEND(incSend, offset);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
__device__ __forceinline__ void loadRecvConn(struct ncclConnInfo* conn, int i) {
|
2020-07-08 11:06:50 -07:00
|
|
|
recvBuff[i] = (union ncclLLFifoLine*)LOAD(conn->buffs+NCCL_PROTO_LL);
|
|
|
|
|
recvStep[i] = LOAD(&conn->step);
|
|
|
|
|
if (wid == i) recvConn = conn;
|
2019-11-19 14:57:39 -08:00
|
|
|
nrecv++;
|
|
|
|
|
}
|
|
|
|
|
__device__ __forceinline__ void loadRecvSync() {
|
|
|
|
|
if (tid >= nthreads-WARP_SIZE && wid < nrecv) {
|
2020-07-08 11:06:50 -07:00
|
|
|
recvConnHeadPtr = LOAD(&recvConn->head);
|
|
|
|
|
recvConnHead = LOAD(&recvConn->step);
|
2019-11-19 14:57:39 -08:00
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
__device__ __forceinline__ void loadSendConn(struct ncclConnInfo* conn, int i) {
|
2020-07-08 11:06:50 -07:00
|
|
|
sendBuff[i] = (union ncclLLFifoLine*)LOAD(conn->buffs+NCCL_PROTO_LL);
|
|
|
|
|
sendStep[i] = LOAD(&conn->step);
|
|
|
|
|
if (wid == i) sendConn = conn;
|
2019-11-19 14:57:39 -08:00
|
|
|
nsend++;
|
|
|
|
|
}
|
|
|
|
|
__device__ __forceinline__ void loadSendSync() {
|
|
|
|
|
if (tid < nsend) {
|
2020-07-08 11:06:50 -07:00
|
|
|
sendConnHeadPtr = LOAD(&sendConn->head);
|
|
|
|
|
sendConnHeadCache = LOAD(sendConnHeadPtr);
|
|
|
|
|
sendConnHead = LOAD(&sendConn->step);
|
2020-12-01 11:33:47 -05:00
|
|
|
sendConnFifoPtr = LOAD(&sendConn->sizesFifo);
|
2019-11-19 14:57:39 -08:00
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
__device__ __forceinline__ void saveRecvSync() {
|
|
|
|
|
if (tid >= nthreads-WARP_SIZE && wid < nrecv) {
|
2020-07-08 11:06:50 -07:00
|
|
|
STORE(&recvConn->step, recvConnHead);
|
2019-11-19 14:57:39 -08:00
|
|
|
__threadfence_block();
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
__device__ __forceinline__ void saveSendSync() {
|
|
|
|
|
if (tid < nsend) {
|
2020-07-08 11:06:50 -07:00
|
|
|
STORE(&sendConn->step, sendConnHead);
|
2019-11-19 14:57:39 -08:00
|
|
|
__threadfence_block();
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
2020-07-08 11:06:50 -07:00
|
|
|
public:
|
|
|
|
|
__device__ __forceinline__
|
2020-07-23 12:08:08 -07:00
|
|
|
ncclLLPrimitives(const int tid, const int nthreads, int* recvPeers, int* sendPeers, int stepLines, struct ncclChannel* channel, struct ncclDevComm* comm)
|
|
|
|
|
: comm(comm), tid(tid), nthreads(nthreads), wid(tid%WARP_SIZE), stepLines(stepLines) {
|
2019-11-19 14:57:39 -08:00
|
|
|
// Make sure step is updated before we read it.
|
|
|
|
|
barrier();
|
|
|
|
|
|
|
|
|
|
for (int i=0; i<NRECV && recvPeers[i] >= 0; i++) loadRecvConn(&channel->devPeers[recvPeers[i]].recv.conn, i);
|
|
|
|
|
for (int i=0; i<NSEND && sendPeers[i] >= 0; i++) loadSendConn(&channel->devPeers[sendPeers[i]].send.conn, i);
|
|
|
|
|
loadRecvSync();
|
|
|
|
|
loadSendSync();
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
__device__ void send(const T* src, int nelem) {
|
|
|
|
|
return LLGenericOp<0, 1, 1, 0>(src, NULL, nelem);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
__device__ void recv(T* dst, int nelem) {
|
|
|
|
|
return LLGenericOp<1, 0, 0, 1>(NULL, dst, nelem);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
__device__ void recvReduceSend(const T* src, int nelem) {
|
|
|
|
|
return LLGenericOp<1, 1, 1, 0>(src, NULL, nelem);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
__device__ void recvReduceCopy(const T* src, T* dst, int nelem) {
|
|
|
|
|
return LLGenericOp<1, 0, 1, 1>(src, dst, nelem);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
__device__ void copySend(const T* src, T* dst, int nelem) {
|
|
|
|
|
return LLGenericOp<0, 1, 1, 1>(src, dst, nelem);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
__device__ void recvCopySend(T* dst, int nelem) {
|
|
|
|
|
return LLGenericOp<1, 1, 0, 1>(NULL, dst, nelem);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
__device__ void recvReduceCopySend(const T* src, T* dst, int nelem) {
|
|
|
|
|
return LLGenericOp<1, 1, 1, 1>(src, dst, nelem);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
__device__ __forceinline__ ~ncclLLPrimitives() {
|
|
|
|
|
// Save steps for the next operation
|
|
|
|
|
saveRecvSync();
|
|
|
|
|
saveSendSync();
|
|
|
|
|
}
|
|
|
|
|
};
|