Merge remote-tracking branch 'nccl/master' into develop
[ROCm/rccl commit: 3a919c1f49]
This commit is contained in:
@@ -8,78 +8,104 @@
|
||||
#ifndef NCCL_COLLECTIVES_H_
|
||||
#define NCCL_COLLECTIVES_H_
|
||||
|
||||
#define FUNC_INDEX_P2P (NCCL_NUM_FUNCTIONS*NCCL_NUM_ALGORITHMS*NCCL_NUM_PROTOCOLS*ncclNumTypes*ncclNumOps)
|
||||
#define FUNC_INDEX(func, redop, ncclType, al, pr) ((((((func)*ncclNumOps + (redop))*ncclNumTypes) + (ncclType))*NCCL_NUM_ALGORITHMS+(al))*NCCL_NUM_PROTOCOLS+(pr))
|
||||
enum ncclDevRedOp_t {
|
||||
ncclDevSum, ncclDevProd, ncclDevMax, ncclDevMin,
|
||||
ncclDevPreMulSum, ncclDevSumPostDiv,
|
||||
ncclNumDevRedOps
|
||||
};
|
||||
struct ncclDevRedOpFull {
|
||||
ncclDevRedOp_t op;
|
||||
bool scalarArgIsPtr;
|
||||
uint64_t scalarArg;
|
||||
};
|
||||
|
||||
#define NCCL_FUNC_NAME(func, algo, proto, redop, type) \
|
||||
ncclFunction_##func##_##algo##_##proto##_##redop##_##type
|
||||
#define FUNC_INDEX_P2P (ncclNumTypes+NCCL_NUM_FUNCTIONS*NCCL_NUM_ALGORITHMS*NCCL_NUM_PROTOCOLS*ncclNumTypes*ncclNumDevRedOps)
|
||||
#define FUNC_INDEX(func, devredop, ncclType, al, pr) ((((((func)*ncclNumDevRedOps + (devredop))*ncclNumTypes) + (ncclType))*NCCL_NUM_ALGORITHMS+(al))*NCCL_NUM_PROTOCOLS+(pr))
|
||||
|
||||
#define NCCL_KERN_NAME(func, algo, proto, redop, type) \
|
||||
ncclKernel_##func##_##algo##_##proto##_##redop##_##type
|
||||
#define NCCL_FUNC_NAME(func, algo, proto, devredop, type) \
|
||||
ncclFunction_##func##_##algo##_##proto##_##devredop##_##type
|
||||
|
||||
#define NCCL_ONERANK_REDUCE_NAME(devredop, type) \
|
||||
ncclFunction_OneRankReduce_##devredop##_##type
|
||||
|
||||
#define NCCL_KERN_NAME(func, algo, proto, devredop, type) \
|
||||
ncclKernel_##func##_##algo##_##proto##_##devredop##_##type
|
||||
|
||||
#define NCCL_IMPL_NAME(func, algo, proto) \
|
||||
nccl##func##algo##proto
|
||||
|
||||
/* Declare all collective operations */
|
||||
#define DECL5(func, algo, proto, redop, type) \
|
||||
extern __device__ __attribute__((noinline)) void NCCL_FUNC_NAME(func, algo, proto, redop, type)(struct ncclWorkElem* args); \
|
||||
extern __global__ void NCCL_KERN_NAME(func, algo, proto, redop, type)(struct ncclWorkElem first); \
|
||||
#define DECL5(func, algo, proto, devredop, type) \
|
||||
extern __device__ __attribute__((noinline)) void NCCL_FUNC_NAME(func, algo, proto, devredop, type)(struct ncclWorkElem* args); \
|
||||
extern __global__ void NCCL_KERN_NAME(func, algo, proto, devredop, type)(ncclWorkElem c); \
|
||||
|
||||
//extern __device__ void NCCL_FUNC_NAME(func, algo, proto, redop, type)(); \
|
||||
//extern __global__ void NCCL_KERN_NAME(func, algo, proto, redop, type)(ncclWorkElem c); \
|
||||
#define CONCAT(a,b) a##b
|
||||
#define MACRO_IF(cond, t, f) CONCAT(MACRO_IF_, cond)(t, f)
|
||||
#define MACRO_IF_0(t, f) f
|
||||
#define MACRO_IF_1(t, f) t
|
||||
|
||||
#define DECL4(func, algo, redop, type) \
|
||||
DECL5(func, algo, SIMPLE, redop, type) \
|
||||
DECL5(func, algo, LL, redop, type) \
|
||||
DECL5(func, algo, LL128, redop, type)
|
||||
#define DECL4(func, algo, devredop, type, undef) \
|
||||
MACRO_IF(undef, /*undefined*/, DECL5(func, algo, SIMPLE, devredop, type)) \
|
||||
MACRO_IF(undef, /*undefined*/, DECL5(func, algo, LL, devredop, type)) \
|
||||
MACRO_IF(undef, /*undefined*/, DECL5(func, algo, LL128, devredop, type))
|
||||
|
||||
#define DECL3(func, redop, type) \
|
||||
DECL4(func, RING, redop, type) \
|
||||
DECL4(func, TREE, redop, type) \
|
||||
DECL4(func, COLLNET, redop, type)
|
||||
#define DECL3(func, devredop, type, undef) \
|
||||
DECL4(func, RING, devredop, type, undef) \
|
||||
DECL4(func, TREE, devredop, type, undef) \
|
||||
DECL4(func, COLLNET, devredop, type, undef)
|
||||
|
||||
#if defined(__CUDA_BF16_TYPES_EXIST__)
|
||||
#define DECL2(func, redop) \
|
||||
DECL3(func, redop, int8_t) \
|
||||
DECL3(func, redop, uint8_t) \
|
||||
DECL3(func, redop, int32_t) \
|
||||
DECL3(func, redop, uint32_t) \
|
||||
DECL3(func, redop, int64_t) \
|
||||
DECL3(func, redop, uint64_t) \
|
||||
DECL3(func, redop, half) \
|
||||
DECL3(func, redop, float) \
|
||||
DECL3(func, redop, double) \
|
||||
DECL3(func, redop, __nv_bfloat16)
|
||||
#if defined(RCCL_BFLOAT16)
|
||||
#define DECL2(func, devredop, undefForFloat) \
|
||||
DECL3(func, devredop, int8_t, /*undef=*/0) \
|
||||
DECL3(func, devredop, uint8_t, /*undef=*/0) \
|
||||
DECL3(func, devredop, int32_t, /*undef=*/0) \
|
||||
DECL3(func, devredop, uint32_t, /*undef=*/0) \
|
||||
DECL3(func, devredop, int64_t, /*undef=*/0) \
|
||||
DECL3(func, devredop, uint64_t, /*undef=*/0) \
|
||||
DECL3(func, devredop, half, /*undef=*/undefForFloat) \
|
||||
DECL3(func, devredop, float, /*undef=*/undefForFloat) \
|
||||
DECL3(func, devredop, double, /*undef=*/undefForFloat) \
|
||||
DECL3(func, devredop, rccl_bfloat16, /*undef=*/undefForFloat)
|
||||
#else
|
||||
#define DECL2(func, redop) \
|
||||
DECL3(func, redop, int8_t) \
|
||||
DECL3(func, redop, uint8_t) \
|
||||
DECL3(func, redop, int32_t) \
|
||||
DECL3(func, redop, uint32_t) \
|
||||
DECL3(func, redop, int64_t) \
|
||||
DECL3(func, redop, uint64_t) \
|
||||
DECL3(func, redop, half) \
|
||||
DECL3(func, redop, float) \
|
||||
DECL3(func, redop, double) \
|
||||
DECL3(func, redop, rccl_bfloat16)
|
||||
#define DECL2(func, devredop, undefForFloat) \
|
||||
DECL3(func, devredop, int8_t, /*undef=*/0) \
|
||||
DECL3(func, devredop, uint8_t, /*undef=*/0) \
|
||||
DECL3(func, devredop, int32_t, /*undef=*/0) \
|
||||
DECL3(func, devredop, uint32_t, /*undef=*/0) \
|
||||
DECL3(func, devredop, int64_t, /*undef=*/0) \
|
||||
DECL3(func, devredop, uint64_t, /*undef=*/0) \
|
||||
DECL3(func, devredop, half, /*undef=*/undefForFloat) \
|
||||
DECL3(func, devredop, float, /*undef=*/undefForFloat) \
|
||||
DECL3(func, devredop, double, /*undef=*/undefForFloat)
|
||||
#endif
|
||||
|
||||
#define DECL(func) \
|
||||
DECL2(func, Sum) \
|
||||
DECL2(func, Prod) \
|
||||
DECL2(func, Min) \
|
||||
DECL2(func, Max) \
|
||||
DECL2(func, Avg) \
|
||||
DECL2(func, Sum, /*undefForFloat=*/0) \
|
||||
DECL2(func, Prod, /*undefForFloat=*/0) \
|
||||
DECL2(func, Min, /*undefForFloat=*/0) \
|
||||
DECL2(func, Max, /*undefForFloat=*/0) \
|
||||
DECL2(func, PreMulSum, /*undefForFloat=*/0) \
|
||||
DECL2(func, SumPostDiv, /*undefForFloat=*/1)
|
||||
|
||||
#define DECL_ALL \
|
||||
DECL2(Broadcast, Sum) \
|
||||
DECL(Reduce) \
|
||||
DECL2(AllGather, Sum) \
|
||||
DECL(ReduceScatter) \
|
||||
DECL(AllReduce) \
|
||||
DECL5(SendRecv, RING, SIMPLE, Sum, int8_t) \
|
||||
DECL2(Broadcast, Sum, /*undefForFloat=*/0)
|
||||
DECL(Reduce)
|
||||
DECL2(AllGather, Sum, /*undefForFloat=*/0)
|
||||
DECL(ReduceScatter)
|
||||
DECL(AllReduce)
|
||||
DECL5(SendRecv, RING, SIMPLE, Sum, int8_t)
|
||||
|
||||
DECL_ALL
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, int8_t)(struct ncclWorkElem* args);
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, uint8_t)(struct ncclWorkElem* args);
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, int32_t)(struct ncclWorkElem* args);
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, uint32_t)(struct ncclWorkElem* args);
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, int64_t)(struct ncclWorkElem* args);
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, uint64_t)(struct ncclWorkElem* args);
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, half)(struct ncclWorkElem* args);
|
||||
#if defined(RCCL_BFLOAT16)
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, rccl_bfloat16)(struct ncclWorkElem* args);
|
||||
#endif
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, float)(struct ncclWorkElem* args);
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, double)(struct ncclWorkElem* args);
|
||||
|
||||
// CHUNKSIZE must be a multiple of SLICESIZE
|
||||
//#define ALLREDUCE_SLICESTEPS (NCCL_STEPS/4)
|
||||
@@ -99,5 +125,6 @@ DECL_ALL
|
||||
#define REDUCE_SLICESTEPS 1
|
||||
#define REDUCE_CHUNKSTEPS 1
|
||||
#define SENDRECV_SLICEFACTOR 1
|
||||
#define NCCL_MAX_SLICE_PER_CHUNK 2 // max value for CHUNKSTEPS/SLICESTEPS, must accord with above
|
||||
|
||||
#endif
|
||||
|
||||
Reference in New Issue
Block a user