[DEV] Configure functions in RCCL (#986)
* configure functions in rccl
[ROCm/rccl commit: 28d9b170c9]
This commit is contained in:
@@ -1,11 +0,0 @@
|
||||
/*************************************************************************
|
||||
* Copyright (c) 2015-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* See LICENSE.txt for license information
|
||||
************************************************************************/
|
||||
|
||||
#include "all_gather.h"
|
||||
#include "common.h"
|
||||
#include "collectives.h"
|
||||
|
||||
IMPL_COLL_C(AllGather);
|
||||
@@ -1,12 +0,0 @@
|
||||
/*************************************************************************
|
||||
* Copyright (c) 2015-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* See LICENSE.txt for license information
|
||||
************************************************************************/
|
||||
/*This file is now generated in CMake*/
|
||||
|
||||
// #include "all_reduce.h"
|
||||
// #include "common.h"
|
||||
// #include "collectives.h"
|
||||
|
||||
// IMPL_COLL_R(AllReduce);
|
||||
@@ -1,11 +0,0 @@
|
||||
/*************************************************************************
|
||||
* Copyright (c) 2015-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* See LICENSE.txt for license information
|
||||
************************************************************************/
|
||||
|
||||
#include "alltoall_pivot.h"
|
||||
#include "common.h"
|
||||
#include "collectives.h"
|
||||
|
||||
IMPL_COLL_F(AllToAllPivot);
|
||||
@@ -1,11 +0,0 @@
|
||||
/*************************************************************************
|
||||
* Copyright (c) 2015-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* See LICENSE.txt for license information
|
||||
************************************************************************/
|
||||
|
||||
#include "broadcast.h"
|
||||
#include "common.h"
|
||||
#include "collectives.h"
|
||||
|
||||
IMPL_COLL_C(Broadcast);
|
||||
@@ -32,243 +32,14 @@
|
||||
{ __atomic_store_n((DST), (SRC), __ATOMIC_SEQ_CST); }
|
||||
#endif
|
||||
|
||||
#if defined(ENABLE_LL128) && defined(__gfx90a__)
|
||||
#define NCCL_FUNC5(func, algo, devredop, type, nullify) \
|
||||
MACRO_IF(nullify, nullptr, NCCL_FUNC_NAME(func, algo, LL, devredop, type)), \
|
||||
MACRO_IF(nullify, nullptr, NCCL_FUNC_NAME(func, algo, LL128, devredop, type)), \
|
||||
MACRO_IF(nullify, nullptr, NCCL_FUNC_NAME(func, algo, SIMPLE, devredop, type))
|
||||
#else
|
||||
#define NCCL_FUNC5(func, algo, devredop, type, nullify) \
|
||||
MACRO_IF(nullify, nullptr, NCCL_FUNC_NAME(func, algo, LL, devredop, type)), \
|
||||
MACRO_IF(nullify, nullptr, NCCL_FUNC_NAME(func, algo, LL, devredop, type)), \
|
||||
MACRO_IF(nullify, nullptr, NCCL_FUNC_NAME(func, algo, SIMPLE, devredop, type))
|
||||
#endif
|
||||
|
||||
#define NCCL_FUNC4(func, devredop, type, nullify) \
|
||||
NCCL_FUNC5(func, TREE, devredop, type, nullify), \
|
||||
NCCL_FUNC5(func, RING, devredop, type, nullify), \
|
||||
NCCL_FUNC5(func, COLLNET_DIRECT, devredop, type, nullify), \
|
||||
NCCL_FUNC5(func, COLLNET_CHAIN, devredop, type, nullify), \
|
||||
NCCL_FUNC5(func, NVLS, devredop, type, nullify), \
|
||||
NCCL_FUNC5(func, NVLS_TREE, devredop, type, nullify)
|
||||
|
||||
// Must be consistent with ncclDataType_t
|
||||
#define NCCL_FUNCS3A(func, devredop, nullForFloat) \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, uint8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int32_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, uint32_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int64_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, uint64_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, half, nullForFloat), \
|
||||
NCCL_FUNC4(func, devredop, float, nullForFloat), \
|
||||
NCCL_FUNC4(func, devredop, double, nullForFloat), \
|
||||
NCCL_FUNC4(func, devredop, rccl_bfloat16, nullForFloat)
|
||||
#define NCCL_FUNCS3B(func, devredop) \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0)
|
||||
|
||||
// Must be consistent with ncclRedOp_t
|
||||
#define NCCL_FUNCS2A(func) \
|
||||
NCCL_FUNCS3A(func, Sum, /*nullForFloat=*/0), \
|
||||
NCCL_FUNCS3A(func, Prod, /*nullForFloat=*/0), \
|
||||
NCCL_FUNCS3A(func, Max, /*nullForFloat=*/0), \
|
||||
NCCL_FUNCS3A(func, Min, /*nullForFloat=*/0), \
|
||||
NCCL_FUNCS3A(func, PreMulSum, /*nullForFloat=*/0), \
|
||||
NCCL_FUNCS3A(func, SumPostDiv, /*nullForFloat=*/1)
|
||||
|
||||
#define NCCL_FUNCS2B(func) \
|
||||
NCCL_FUNCS3B(func, Sum), \
|
||||
NCCL_FUNCS3B(func, Sum), \
|
||||
NCCL_FUNCS3B(func, Sum), \
|
||||
NCCL_FUNCS3B(func, Sum), \
|
||||
NCCL_FUNCS3B(func, Sum), \
|
||||
NCCL_FUNCS3B(func, Sum)
|
||||
|
||||
// Must be consistent with the ncclFuncSet enum
|
||||
using ncclKernelFunc_t = void (*)();
|
||||
|
||||
static const __device__ constexpr ncclKernelFunc_t ncclFuncs[]{
|
||||
// Don't try to initialize the host shadow copy of this device-side global
|
||||
// variable. There is no host pointer to a device-side function, which
|
||||
// confuses clang. This will be fixed in the next clang release.
|
||||
#if defined(__HIP_DEVICE_COMPILE__)
|
||||
#if defined(BUILD_ALLREDUCE_ONLY)
|
||||
NCCL_FUNC4(AllReduce, Sum, float, 0),
|
||||
#else
|
||||
NCCL_FUNCS2B(Broadcast),
|
||||
NCCL_FUNCS2A(Reduce),
|
||||
NCCL_FUNCS2B(AllGather),
|
||||
NCCL_FUNCS2A(ReduceScatter),
|
||||
NCCL_FUNCS2A(AllReduce),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, int8_t),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, uint8_t),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, int32_t),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, uint32_t),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, int64_t),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, uint64_t),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, half),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, float),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, double),
|
||||
#if defined(RCCL_BFLOAT16)
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, rccl_bfloat16),
|
||||
#endif
|
||||
NCCL_FUNC_NAME(SendRecv, RING, SIMPLE, Sum, int8_t),
|
||||
NCCL_FUNC_NAME(AllToAllPivot, RING, SIMPLE, Sum, int8_t),
|
||||
#endif
|
||||
#endif
|
||||
};
|
||||
// Defined in device_table.cpp
|
||||
extern __device__ ncclKernelFunc_t const ncclFuncs[];
|
||||
|
||||
static_assert(FUNC_INDEX_P2P == 5410, "Wrong P2P function index");
|
||||
static_assert(FUNC_INDEX_ALLTOALL_PIVOT == 5411, "Wrong AllToAllPivot function index");
|
||||
|
||||
#if !defined(USE_INDIRECT_FUNCTION_CALL) || defined(__gfx940__) || defined(__gfx941__) || defined(__gfx942__)
|
||||
template<unsigned short f, unsigned short l, bool u>
|
||||
struct Caller {
|
||||
static __forceinline__ __device__ __host__
|
||||
void call(unsigned short funcIndex) noexcept
|
||||
{
|
||||
constexpr unsigned short m = f + (l - f) / 2;
|
||||
|
||||
return (funcIndex < m) ? Caller<f, m, u>::call(funcIndex) : Caller<m, l, u>::call(funcIndex);
|
||||
}
|
||||
};
|
||||
|
||||
template<unsigned short f, bool u>
|
||||
struct Caller<f, f + 1, u>{
|
||||
static __forceinline__ __device__ __host__
|
||||
void call(unsigned short funcIndex) noexcept { ncclFuncs[f](); }
|
||||
};
|
||||
|
||||
template<bool USING_LL128>
|
||||
__forceinline__
|
||||
__device__
|
||||
void NCCL_CALL_FUNCTIONS(unsigned short funcIndex) noexcept {
|
||||
#if defined(BUILD_ALLREDUCE_ONLY)
|
||||
if (funcIndex == FUNC_INDEX(ncclFuncAllReduce, ncclSum, ncclFloat32, NCCL_ALGO_RING, NCCL_PROTO_SIMPLE))
|
||||
ncclFunction_AllReduce_RING_SIMPLE_Sum_float();
|
||||
else if (funcIndex == FUNC_INDEX(ncclFuncAllReduce, ncclSum, ncclFloat32, NCCL_ALGO_RING, NCCL_PROTO_LL))
|
||||
ncclFunction_AllReduce_RING_LL_Sum_float();
|
||||
else if (USING_LL128 && funcIndex == FUNC_INDEX(ncclFuncAllReduce, ncclSum, ncclFloat32, NCCL_ALGO_RING, NCCL_PROTO_LL128))
|
||||
ncclFunction_AllReduce_RING_LL128_Sum_float();
|
||||
else if (!USING_LL128 && funcIndex == FUNC_INDEX(ncclFuncAllReduce, ncclSum, ncclFloat32, NCCL_ALGO_RING, NCCL_PROTO_LL128))
|
||||
ncclFunction_AllReduce_RING_LL_Sum_float();
|
||||
else if (funcIndex == FUNC_INDEX(ncclFuncAllReduce, ncclSum, ncclFloat32, NCCL_ALGO_TREE, NCCL_PROTO_SIMPLE))
|
||||
ncclFunction_AllReduce_TREE_SIMPLE_Sum_float();
|
||||
else if (funcIndex == FUNC_INDEX(ncclFuncAllReduce, ncclSum, ncclFloat32, NCCL_ALGO_TREE, NCCL_PROTO_LL))
|
||||
ncclFunction_AllReduce_TREE_LL_Sum_float();
|
||||
else if (USING_LL128 && funcIndex == FUNC_INDEX(ncclFuncAllReduce, ncclSum, ncclFloat32, NCCL_ALGO_TREE, NCCL_PROTO_LL128))
|
||||
ncclFunction_AllReduce_TREE_LL128_Sum_float();
|
||||
else if (!USING_LL128 && funcIndex == FUNC_INDEX(ncclFuncAllReduce, ncclSum, ncclFloat32, NCCL_ALGO_TREE, NCCL_PROTO_LL128))
|
||||
ncclFunction_AllReduce_TREE_LL_Sum_float();
|
||||
else if (funcIndex == FUNC_INDEX(ncclFuncAllReduce, ncclSum, ncclFloat32, NCCL_ALGO_COLLNET_DIRECT, NCCL_PROTO_SIMPLE))
|
||||
ncclFunction_AllReduce_COLLNET_DIRECT_SIMPLE_Sum_float();
|
||||
else if (funcIndex == FUNC_INDEX(ncclFuncAllReduce, ncclSum, ncclFloat32, NCCL_ALGO_COLLNET_DIRECT, NCCL_PROTO_LL))
|
||||
ncclFunction_AllReduce_COLLNET_DIRECT_LL_Sum_float();
|
||||
else if (USING_LL128 && funcIndex == FUNC_INDEX(ncclFuncAllReduce, ncclSum, ncclFloat32, NCCL_ALGO_COLLNET_DIRECT, NCCL_PROTO_LL128))
|
||||
ncclFunction_AllReduce_COLLNET_DIRECT_LL128_Sum_float();
|
||||
else if (!USING_LL128 && funcIndex == FUNC_INDEX(ncclFuncAllReduce, ncclSum, ncclFloat32, NCCL_ALGO_COLLNET_DIRECT, NCCL_PROTO_LL128))
|
||||
ncclFunction_AllReduce_COLLNET_DIRECT_LL_Sum_float();
|
||||
else if (funcIndex == FUNC_INDEX(ncclFuncAllReduce, ncclSum, ncclFloat32, NCCL_ALGO_COLLNET_CHAIN, NCCL_PROTO_SIMPLE))
|
||||
ncclFunction_AllReduce_COLLNET_CHAIN_SIMPLE_Sum_float();
|
||||
else if (funcIndex == FUNC_INDEX(ncclFuncAllReduce, ncclSum, ncclFloat32, NCCL_ALGO_COLLNET_CHAIN, NCCL_PROTO_LL))
|
||||
ncclFunction_AllReduce_COLLNET_CHAIN_LL_Sum_float();
|
||||
else if (USING_LL128 && funcIndex == FUNC_INDEX(ncclFuncAllReduce, ncclSum, ncclFloat32, NCCL_ALGO_COLLNET_CHAIN, NCCL_PROTO_LL128))
|
||||
ncclFunction_AllReduce_COLLNET_CHAIN_LL128_Sum_float();
|
||||
else if (!USING_LL128 && funcIndex == FUNC_INDEX(ncclFuncAllReduce, ncclSum, ncclFloat32, NCCL_ALGO_COLLNET_CHAIN, NCCL_PROTO_LL128))
|
||||
ncclFunction_AllReduce_COLLNET_CHAIN_LL_Sum_float();
|
||||
else
|
||||
assert("Unsupported function index");
|
||||
#else
|
||||
if (funcIndex < 1080) {
|
||||
if (funcIndex % 18 == 0) ncclFunction_Broadcast_TREE_LL_Sum_int8_t();
|
||||
else if (USING_LL128 && funcIndex % 18 == 1) ncclFunction_Broadcast_TREE_LL128_Sum_int8_t();
|
||||
else if (!USING_LL128 && funcIndex % 18 == 1) ncclFunction_Broadcast_TREE_LL_Sum_int8_t();
|
||||
else if (funcIndex % 18 == 2) ncclFunction_Broadcast_TREE_SIMPLE_Sum_int8_t();
|
||||
else if (funcIndex % 18 == 3) ncclFunction_Broadcast_RING_LL_Sum_int8_t();
|
||||
else if (USING_LL128 && funcIndex % 18 == 4) ncclFunction_Broadcast_RING_LL128_Sum_int8_t();
|
||||
else if (!USING_LL128 && funcIndex % 18 == 4) ncclFunction_Broadcast_RING_LL_Sum_int8_t();
|
||||
else if (funcIndex % 18 == 5) ncclFunction_Broadcast_RING_SIMPLE_Sum_int8_t();
|
||||
else if (funcIndex % 18 == 6) ncclFunction_Broadcast_COLLNET_DIRECT_LL_Sum_int8_t();
|
||||
else if (USING_LL128 && funcIndex % 18 == 7) ncclFunction_Broadcast_COLLNET_DIRECT_LL128_Sum_int8_t();
|
||||
else if (!USING_LL128 && funcIndex % 18 == 7) ncclFunction_Broadcast_COLLNET_DIRECT_LL_Sum_int8_t();
|
||||
else if (funcIndex % 18 == 8) ncclFunction_Broadcast_COLLNET_DIRECT_SIMPLE_Sum_int8_t();
|
||||
else if (funcIndex % 18 == 9) ncclFunction_Broadcast_COLLNET_CHAIN_LL_Sum_int8_t();
|
||||
else if (USING_LL128 && funcIndex % 18 == 10) ncclFunction_Broadcast_COLLNET_CHAIN_LL128_Sum_int8_t();
|
||||
else if (!USING_LL128 && funcIndex % 18 == 10) ncclFunction_Broadcast_COLLNET_CHAIN_LL_Sum_int8_t();
|
||||
else ncclFunction_Broadcast_COLLNET_CHAIN_SIMPLE_Sum_int8_t();
|
||||
}
|
||||
else if (funcIndex < 2160) Caller<1080, 2160, USING_LL128>::call(funcIndex);
|
||||
else if (funcIndex < 3240) {
|
||||
if (funcIndex % 18 == 0) ncclFunction_AllGather_TREE_LL_Sum_int8_t();
|
||||
else if (USING_LL128 && funcIndex % 18 == 1) ncclFunction_AllGather_TREE_LL128_Sum_int8_t();
|
||||
else if (!USING_LL128 && funcIndex % 18 == 1) ncclFunction_AllGather_TREE_LL_Sum_int8_t();
|
||||
else if (funcIndex % 18 == 2) ncclFunction_AllGather_TREE_SIMPLE_Sum_int8_t();
|
||||
else if (funcIndex % 18 == 3) ncclFunction_AllGather_RING_LL_Sum_int8_t();
|
||||
else if (USING_LL128 && funcIndex % 18 == 4) ncclFunction_AllGather_RING_LL128_Sum_int8_t();
|
||||
else if (!USING_LL128 && funcIndex % 18 == 4) ncclFunction_AllGather_RING_LL_Sum_int8_t();
|
||||
else if (funcIndex % 18 == 5) ncclFunction_AllGather_RING_SIMPLE_Sum_int8_t();
|
||||
else if (funcIndex % 18 == 6) ncclFunction_AllGather_COLLNET_DIRECT_LL_Sum_int8_t();
|
||||
else if (USING_LL128 && funcIndex % 18 == 7) ncclFunction_AllGather_COLLNET_DIRECT_LL128_Sum_int8_t();
|
||||
else if (!USING_LL128 && funcIndex % 18 == 7) ncclFunction_AllGather_COLLNET_DIRECT_LL_Sum_int8_t();
|
||||
else if (funcIndex % 18 == 8) ncclFunction_AllGather_COLLNET_DIRECT_SIMPLE_Sum_int8_t();
|
||||
else if (funcIndex % 18 == 9) ncclFunction_AllGather_COLLNET_CHAIN_LL_Sum_int8_t();
|
||||
else if (USING_LL128 && funcIndex % 18 == 10) ncclFunction_AllGather_COLLNET_CHAIN_LL128_Sum_int8_t();
|
||||
else if (!USING_LL128 && funcIndex % 18 == 10) ncclFunction_AllGather_COLLNET_CHAIN_LL_Sum_int8_t();
|
||||
else ncclFunction_AllGather_COLLNET_CHAIN_SIMPLE_Sum_int8_t();
|
||||
}
|
||||
else if (funcIndex < 5400) Caller<3240, 5400, USING_LL128>::call(funcIndex);
|
||||
else {
|
||||
switch (funcIndex - 5400) {
|
||||
case 0:
|
||||
ncclFunction_OneRankReduce_PreMulSum_int8_t();
|
||||
break;
|
||||
case 1:
|
||||
ncclFunction_OneRankReduce_PreMulSum_uint8_t();
|
||||
break;
|
||||
case 2:
|
||||
ncclFunction_OneRankReduce_PreMulSum_int32_t();
|
||||
break;
|
||||
case 3:
|
||||
ncclFunction_OneRankReduce_PreMulSum_uint32_t();
|
||||
break;
|
||||
case 4:
|
||||
ncclFunction_OneRankReduce_PreMulSum_int64_t();
|
||||
break;
|
||||
case 5:
|
||||
ncclFunction_OneRankReduce_PreMulSum_uint64_t();
|
||||
break;
|
||||
case 6:
|
||||
ncclFunction_OneRankReduce_PreMulSum_half();
|
||||
break;
|
||||
case 7:
|
||||
ncclFunction_OneRankReduce_PreMulSum_float();
|
||||
break;
|
||||
case 8:
|
||||
ncclFunction_OneRankReduce_PreMulSum_double();
|
||||
break;
|
||||
case 9:
|
||||
ncclFunction_OneRankReduce_PreMulSum_rccl_bfloat16();
|
||||
break;
|
||||
case 10:
|
||||
ncclFunction_SendRecv_RING_SIMPLE_Sum_int8_t();
|
||||
break;
|
||||
case 11:
|
||||
ncclFunction_AllToAllPivot_RING_SIMPLE_Sum_int8_t();
|
||||
default:
|
||||
break;
|
||||
}
|
||||
}
|
||||
#endif
|
||||
}
|
||||
#ifndef USE_INDIRECT_FUNCTION_CALL
|
||||
__device__ void NCCL_CALL_FUNCTIONS(unsigned short funcIndex) noexcept;
|
||||
#endif
|
||||
|
||||
template <ncclFunc_t FUNCTION, int ALGO, int PROTO, class REDOP, typename T, int UNROLL>
|
||||
@@ -464,7 +235,7 @@ static __forceinline__ __device__ void ncclRedopPtrDeref(struct ncclWorkElem* we
|
||||
}
|
||||
}
|
||||
|
||||
template<ncclFunc_t Fn, typename T, typename RedOp, int Algo, int Proto, int FnIndex, bool COLLTRACE>
|
||||
template<bool COLLTRACE>
|
||||
__forceinline__ __device__ void ncclKernel(
|
||||
struct ncclDevComm* comm, uint64_t channelMask, struct ncclWork* workHead
|
||||
) {
|
||||
@@ -565,19 +336,12 @@ __forceinline__ __device__ void ncclKernel(
|
||||
__synclds();
|
||||
|
||||
if (tid == 0) __insert_timestamp(__LINE__);
|
||||
if (ncclShmem.work.header.funcIndex == FnIndex) {
|
||||
RunWork<Fn, T, RedOp, Algo, Proto>().run(&ncclShmem.work);
|
||||
} else {
|
||||
#if defined(USE_INDIRECT_FUNCTION_CALL) && !defined(__gfx940__) && !defined(__gfx941__) && !defined(__gfx942__)
|
||||
ncclFuncs[ncclShmem.work.header.funcIndex]();
|
||||
|
||||
#ifdef USE_INDIRECT_FUNCTION_CALL
|
||||
ncclFuncs[ncclShmem.work.header.funcIndex]();
|
||||
#else
|
||||
#if defined(ENABLE_LL128) && defined(__gfx90a__)
|
||||
NCCL_CALL_FUNCTIONS<1>(ncclShmem.work.header.funcIndex);
|
||||
#else
|
||||
NCCL_CALL_FUNCTIONS<0>(ncclShmem.work.header.funcIndex);
|
||||
NCCL_CALL_FUNCTIONS(ncclShmem.work.header.funcIndex);
|
||||
#endif
|
||||
#endif
|
||||
}
|
||||
|
||||
int workIxNext = ncclShmem.work.header.workNext;
|
||||
__synclds();
|
||||
@@ -606,28 +370,25 @@ __forceinline__ __device__ void ncclKernel(
|
||||
}
|
||||
|
||||
#ifdef ENABLE_COLLTRACE
|
||||
#define IMPL_COLL_KERN(func, algo, proto, devredop, type, fIndex) \
|
||||
#define IMPL_MAIN_KERN() \
|
||||
__launch_bounds__(NCCL_MAX_NTHREADS, 1) \
|
||||
__global__ void NCCL_KERN_NAME(func, algo, proto, devredop, type)(struct ncclDevComm* comm, uint64_t channelMask, struct ncclWork* workHead) { \
|
||||
ncclKernel<ncclFunc##func, type, Func##devredop<type>, NCCL_ALGO_##algo, NCCL_PROTO_##proto, fIndex, false>(comm, channelMask, workHead); \
|
||||
__global__ void rccl_main_kernel(struct ncclDevComm* comm, uint64_t channelMask, struct ncclWork* workHead) { \
|
||||
ncclKernel<false>(comm, channelMask, workHead); \
|
||||
} \
|
||||
\
|
||||
__launch_bounds__(NCCL_MAX_NTHREADS, 1) \
|
||||
__global__ void NCCL_KERN_NAME_DEBUG(func, algo, proto, devredop, type)(struct ncclDevComm* comm, uint64_t channelMask, struct ncclWork* workHead) { \
|
||||
ncclKernel<ncclFunc##func, type, Func##devredop<type>, NCCL_ALGO_##algo, NCCL_PROTO_##proto, fIndex, true>(comm, channelMask, workHead); \
|
||||
__global__ void rccl_main_kernel_debug(struct ncclDevComm* comm, uint64_t channelMask, struct ncclWork* workHead) { \
|
||||
ncclKernel<true>(comm, channelMask, workHead); \
|
||||
}
|
||||
#else
|
||||
#define IMPL_COLL_KERN(func, algo, proto, devredop, type, fIndex) \
|
||||
#define IMPL_MAIN_KERN() \
|
||||
__launch_bounds__(NCCL_MAX_NTHREADS, 1) \
|
||||
__global__ void NCCL_KERN_NAME(func, algo, proto, devredop, type)(struct ncclDevComm* comm, uint64_t channelMask, struct ncclWork* workHead) { \
|
||||
ncclKernel<ncclFunc##func, type, Func##devredop<type>, NCCL_ALGO_##algo, NCCL_PROTO_##proto, fIndex, false>(comm, channelMask, workHead); \
|
||||
__global__ void rccl_main_kernel(struct ncclDevComm* comm, uint64_t channelMask, struct ncclWork* workHead) { \
|
||||
ncclKernel<false>(comm, channelMask, workHead); \
|
||||
}
|
||||
#endif
|
||||
|
||||
// Examples : AllReduce, RING, LL, Sum, uint8
|
||||
/* Functions for aggregation case */
|
||||
|
||||
#if defined(USE_INDIRECT_FUNCTION_CALL) && !defined(__gfx940__) && !defined(__gfx941__) && !defined(__gfx942__)
|
||||
#ifdef USE_INDIRECT_FUNCTION_CALL
|
||||
#define IMPL_COLL_FUNC(func, algo, proto, devredop, type) \
|
||||
__device__ void NCCL_FUNC_NAME(func, algo, proto, devredop, type)() { \
|
||||
RunWork<ncclFunc##func, type, Func##devredop<type>, NCCL_ALGO_##algo, NCCL_PROTO_##proto>().run(&ncclShmem.work); \
|
||||
@@ -639,67 +400,6 @@ __device__ __attribute__((noinline)) void NCCL_FUNC_NAME(func, algo, proto, dev
|
||||
}
|
||||
#endif
|
||||
|
||||
// Only generate inline kernels for LL
|
||||
#if defined(ENABLE_LL128) && defined(__gfx90a__)
|
||||
#define IMPL_COLL4(func, algo, devredop, type) \
|
||||
IMPL_COLL_FUNC(func, algo, LL, devredop, type) \
|
||||
IMPL_COLL_FUNC(func, algo, LL128, devredop, type) \
|
||||
IMPL_COLL_FUNC(func, algo, SIMPLE, devredop, type)
|
||||
#else
|
||||
#define IMPL_COLL4(func, algo, devredop, type) \
|
||||
IMPL_COLL_FUNC(func, algo, LL, devredop, type) \
|
||||
IMPL_COLL_FUNC(func, algo, SIMPLE, devredop, type)
|
||||
#endif
|
||||
|
||||
#define IMPL_COLL3(func, devredop, type) \
|
||||
IMPL_COLL4(func, TREE, devredop, type) \
|
||||
IMPL_COLL4(func, RING, devredop, type) \
|
||||
IMPL_COLL4(func, COLLNET_DIRECT, devredop, type) \
|
||||
IMPL_COLL4(func, COLLNET_CHAIN, devredop, type) \
|
||||
IMPL_COLL4(func, NVLS, devredop, type) \
|
||||
IMPL_COLL4(func, NVLS_TREE, devredop, type)
|
||||
|
||||
#define IMPL_COLL2(func, devredop) \
|
||||
IMPL_COLL3(func, devredop, int8_t) \
|
||||
IMPL_COLL3(func, devredop, uint8_t) \
|
||||
IMPL_COLL3(func, devredop, int32_t) \
|
||||
IMPL_COLL3(func, devredop, uint32_t) \
|
||||
IMPL_COLL3(func, devredop, int64_t) \
|
||||
IMPL_COLL3(func, devredop, uint64_t) \
|
||||
IMPL_COLL3(func, devredop, half) \
|
||||
IMPL_COLL3(func, devredop, float) \
|
||||
IMPL_COLL3(func, devredop, double) \
|
||||
IMPL_COLL3(func, devredop, rccl_bfloat16)
|
||||
|
||||
#define IMPL_COLL2A(func, devredop) \
|
||||
IMPL_COLL3(func, devredop, int8_t) \
|
||||
IMPL_COLL3(func, devredop, uint8_t) \
|
||||
IMPL_COLL3(func, devredop, int32_t) \
|
||||
IMPL_COLL3(func, devredop, uint32_t) \
|
||||
IMPL_COLL3(func, devredop, int64_t) \
|
||||
IMPL_COLL3(func, devredop, uint64_t)
|
||||
|
||||
// Reduction define all functions
|
||||
#define IMPL_COLL_R(func) \
|
||||
IMPL_COLL2(func, Sum) \
|
||||
IMPL_COLL2(func, Prod) \
|
||||
IMPL_COLL2(func, Min) \
|
||||
IMPL_COLL2(func, Max) \
|
||||
IMPL_COLL2(func, PreMulSum) \
|
||||
IMPL_COLL2A(func, SumPostDiv)
|
||||
|
||||
// Copy primitives only define one function for copy
|
||||
#define IMPL_COLL_C(func) IMPL_COLL3(func, Sum, int8_t);
|
||||
|
||||
// Point-to-point primitives only have one function/kernel.
|
||||
#define IMPL_COLL_P(func) \
|
||||
IMPL_COLL_FUNC(func, RING, SIMPLE, Sum, int8_t); \
|
||||
IMPL_COLL_KERN(func, RING, SIMPLE, Sum, int8_t, FUNC_INDEX_P2P);
|
||||
|
||||
// AllToAll Pivot primitive only has one function.
|
||||
#define IMPL_COLL_F(func) \
|
||||
IMPL_COLL_FUNC(func, RING, SIMPLE, Sum, int8_t);
|
||||
|
||||
#define NCCL_NVLS_ENABLED (__CUDA_ARCH__ >= 900 && NCCL_NVLS_SUPPORTS(NCCL_TYPE, NCCL_OP))
|
||||
|
||||
#endif
|
||||
#endif
|
||||
@@ -1,126 +0,0 @@
|
||||
/*************************************************************************
|
||||
* Copyright (c) 2015-2021, NVIDIA CORPORATION. All rights reserved.
|
||||
* Modifications Copyright (c) 2019-2021 Advanced Micro Devices, Inc. All rights reserved.
|
||||
*
|
||||
* See LICENSE.txt for license information
|
||||
************************************************************************/
|
||||
|
||||
#include "devcomm.h"
|
||||
#include "collectives.h"
|
||||
#include "common.h"
|
||||
|
||||
__shared__ ncclShmemData ncclShmem;
|
||||
#if __CUDA_ARCH__ < 700
|
||||
__shared__ ulong2 ncclShmemPerWarp[ncclShmemScratchWarpSize()*(NCCL_MAX_NTHREADS/WARP_SIZE)/sizeof(ulong2)];
|
||||
#endif
|
||||
|
||||
#if defined(__HIP_PLATFORM_HCC__) || defined(__HCC__) || defined(__HIPCC__)
|
||||
#else
|
||||
#define NCCL_FUNC5(func, algo, devredop, type, nullify) \
|
||||
MACRO_IF(nullify, nullptr, NCCL_FUNC_NAME(func, algo, LL, devredop, type)), \
|
||||
MACRO_IF(nullify, nullptr, NCCL_FUNC_NAME(func, algo, LL128, devredop, type)), \
|
||||
MACRO_IF(nullify, nullptr, NCCL_FUNC_NAME(func, algo, SIMPLE, devredop, type))
|
||||
|
||||
#define NCCL_FUNC4(func, devredop, type, nullify) \
|
||||
NCCL_FUNC5(func, TREE, devredop, type, nullify), \
|
||||
NCCL_FUNC5(func, RING, devredop, type, nullify), \
|
||||
NCCL_FUNC5(func, COLLNET_DIRECT, devredop, type, nullify), \
|
||||
NCCL_FUNC5(func, COLLNET_CHAIN, devredop, type, nullify), \
|
||||
NCCL_FUNC5(func, NVLS, devredop, type, nullify), \
|
||||
NCCL_FUNC5(func, NVLS_TREE, devredop, type, nullify)
|
||||
|
||||
#if defined(__CUDA_BF16_TYPES_EXIST__)
|
||||
// Must be consistent with ncclDataType_t
|
||||
#define NCCL_FUNCS3A(func, devredop, nullForFloat) \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, uint8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int32_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, uint32_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int64_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, uint64_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, half, nullForFloat), \
|
||||
NCCL_FUNC4(func, devredop, float, nullForFloat), \
|
||||
NCCL_FUNC4(func, devredop, double, nullForFloat), \
|
||||
NCCL_FUNC4(func, devredop, __nv_bfloat16, nullForFloat)
|
||||
#define NCCL_FUNCS3B(func, devredop) \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0)
|
||||
#else
|
||||
// Must be consistent with ncclDataType_t
|
||||
#define NCCL_FUNCS3A(func, devredop, nullForFloat) \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, uint8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int32_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, uint32_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int64_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, uint64_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, half, nullForFloat), \
|
||||
NCCL_FUNC4(func, devredop, float, nullForFloat), \
|
||||
NCCL_FUNC4(func, devredop, double, nullForFloat)
|
||||
#define NCCL_FUNCS3B(func, devredop) \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0), \
|
||||
NCCL_FUNC4(func, devredop, int8_t, 0)
|
||||
#endif
|
||||
|
||||
// Must be consistent with ncclRedOp_t
|
||||
#define NCCL_FUNCS2A(func) \
|
||||
NCCL_FUNCS3A(func, Sum, /*nullForFloat=*/0), \
|
||||
NCCL_FUNCS3A(func, Prod, /*nullForFloat=*/0), \
|
||||
NCCL_FUNCS3A(func, Max, /*nullForFloat=*/0), \
|
||||
NCCL_FUNCS3A(func, Min, /*nullForFloat=*/0), \
|
||||
NCCL_FUNCS3A(func, PreMulSum, /*nullForFloat=*/0), \
|
||||
NCCL_FUNCS3A(func, SumPostDiv, /*nullForFloat=*/1)
|
||||
|
||||
#define NCCL_FUNCS2B(func) \
|
||||
NCCL_FUNCS3B(func, Sum), \
|
||||
NCCL_FUNCS3B(func, Sum), \
|
||||
NCCL_FUNCS3B(func, Sum), \
|
||||
NCCL_FUNCS3B(func, Sum), \
|
||||
NCCL_FUNCS3B(func, Sum), \
|
||||
NCCL_FUNCS3B(func, Sum)
|
||||
|
||||
// Must be consistent with the ncclFuncSet enum
|
||||
__device__ ncclKern_t ncclFuncs[1+ncclNumTypes+NCCL_NUM_FUNCTIONS*ncclNumDevRedOps*ncclNumTypes*NCCL_NUM_ALGORITHMS*NCCL_NUM_PROTOCOLS] = {
|
||||
// Don't try to initialize the host shadow copy of this device-side global
|
||||
// variable. There is no host pointer to a device-side function, which
|
||||
// confuses clang. This will be fixed in the next clang release.
|
||||
#if __CUDA_ARCH__
|
||||
NCCL_FUNC_NAME(SendRecv, RING, SIMPLE, Sum, int8_t),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, int8_t),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, uint8_t),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, int32_t),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, uint32_t),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, int64_t),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, uint64_t),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, half),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, float),
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, double),
|
||||
#if defined(__CUDA_BF16_TYPES_EXIST__)
|
||||
NCCL_ONERANK_REDUCE_NAME(PreMulSum, __nv_bfloat16),
|
||||
#endif
|
||||
NCCL_FUNCS2B(Broadcast),
|
||||
NCCL_FUNCS2A(Reduce),
|
||||
NCCL_FUNCS2B(AllGather),
|
||||
NCCL_FUNCS2A(ReduceScatter),
|
||||
NCCL_FUNCS2A(AllReduce)
|
||||
#endif
|
||||
};
|
||||
#endif
|
||||
|
||||
// Workaround for https://reviews.llvm.org/D55580
|
||||
__device__ void ncclWorkaroundClangD55580() {}
|
||||
@@ -1,13 +0,0 @@
|
||||
/*************************************************************************
|
||||
* Copyright (c) 2015-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* See LICENSE.txt for license information
|
||||
************************************************************************/
|
||||
|
||||
/*This file is now generated in CMake*/
|
||||
|
||||
// #include "reduce.h"
|
||||
// #include "common.h"
|
||||
// #include "collectives.h"
|
||||
|
||||
// IMPL_COLL_R(Reduce);
|
||||
@@ -1,13 +0,0 @@
|
||||
/*************************************************************************
|
||||
* Copyright (c) 2015-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* See LICENSE.txt for license information
|
||||
************************************************************************/
|
||||
|
||||
/*This file is now generated in CMake*/
|
||||
|
||||
// #include "reduce_scatter.h"
|
||||
// #include "common.h"
|
||||
// #include "collectives.h"
|
||||
|
||||
// IMPL_COLL_R(ReduceScatter);
|
||||
@@ -1,11 +0,0 @@
|
||||
/*************************************************************************
|
||||
* Copyright (c) 2015-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* See LICENSE.txt for license information
|
||||
************************************************************************/
|
||||
|
||||
#include "sendrecv.h"
|
||||
#include "common.h"
|
||||
#include "collectives.h"
|
||||
|
||||
IMPL_COLL_P(SendRecv);
|
||||
@@ -20,8 +20,6 @@
|
||||
#include <cstring> // std::memcpy
|
||||
#include <cinttypes> // PRIx64
|
||||
|
||||
static void* const ncclKernelGeneric = (void*)NCCL_KERN_NAME(SendRecv, RING, SIMPLE, Sum, int8_t);
|
||||
|
||||
struct ncclKernelMatch {
|
||||
void* kernelFn;
|
||||
bool specialized;
|
||||
@@ -29,15 +27,18 @@ struct ncclKernelMatch {
|
||||
|
||||
typedef void(*ncclKern_t)(struct ncclDevComm* comm, uint64_t channelMask, struct ncclWork* workHead);
|
||||
|
||||
// Definition of rccl_main_kernel which is only used in here
|
||||
IMPL_MAIN_KERN();
|
||||
|
||||
// Must be consistent with the ncclFuncSet enum
|
||||
#ifdef ENABLE_COLLTRACE
|
||||
static ncclKernelMatch const ncclKerns[2] = {
|
||||
{(void *)NCCL_KERN_NAME(SendRecv, RING, SIMPLE, Sum, int8_t), true},
|
||||
{(void *)NCCL_KERN_NAME_DEBUG(SendRecv, RING, SIMPLE, Sum, int8_t), true},
|
||||
{(void *)rccl_main_kernel, true},
|
||||
{(void *)rccl_main_kernel_debug, true},
|
||||
};
|
||||
#else
|
||||
static ncclKernelMatch const ncclKerns[1] = {
|
||||
{(void*)NCCL_KERN_NAME(SendRecv, RING, SIMPLE, Sum, int8_t), true}
|
||||
{(void*)rccl_main_kernel, true}
|
||||
};
|
||||
#endif
|
||||
|
||||
@@ -169,7 +170,7 @@ static void appendWorkElemP2p(
|
||||
struct ncclComm* comm, struct ncclKernelPlan* plan, int channelId,
|
||||
struct ncclWorkElemP2p const *elem, bool fuseOk
|
||||
) {
|
||||
constexpr int funcIndex = FUNC_INDEX_P2P;
|
||||
int funcIndex = ncclFuncId_P2p();
|
||||
struct ncclKernelPlan::Channel* chan = &plan->channels[channelId];
|
||||
struct ncclWorkList* q = ncclIntruQueueTail(&chan->workQueue);
|
||||
if (q && funcIndex == q->work.header.funcIndex) {
|
||||
@@ -191,7 +192,7 @@ static void appendWorkElemP2p(
|
||||
}
|
||||
q = ncclMemoryStackAlloc<struct ncclWorkList>(&comm->memScoped);
|
||||
q->work.header.type = ncclWorkTypeP2p;
|
||||
q->work.header.funcIndex = FUNC_INDEX_P2P;
|
||||
q->work.header.funcIndex = ncclFuncId_P2p();
|
||||
chan->p2pTailElem[ncclWorkP2pTypeRecv-1] = 0;
|
||||
chan->p2pTailElem[ncclWorkP2pTypeSend-1] = 1;
|
||||
q->work.p2pElems[chan->p2pTailElem[elem->p2pType-1]] = *elem; // C++ struct assignment
|
||||
@@ -1313,12 +1314,12 @@ comp_next:
|
||||
|
||||
if (info->comm->nRanks == 1) {
|
||||
// one-rank reduce index
|
||||
*workFuncIndex = FUNC_INDEX_P2P - ncclNumTypes + int(info->datatype);
|
||||
*workFuncIndex = ncclFuncId_P2p() + int(info->datatype);
|
||||
return ncclSuccess;
|
||||
} else if (info->coll == ncclFuncAllToAllPivot) {
|
||||
*workFuncIndex = FUNC_INDEX_ALLTOALL_PIVOT;
|
||||
*workFuncIndex = ncclFuncId_AllToAllPivot();
|
||||
} else {
|
||||
*workFuncIndex = FUNC_INDEX(info->coll, info->opFull.op, info->datatype, info->algorithm, info->protocol);
|
||||
*workFuncIndex = ncclFuncId(info->coll, info->opFull.op, info->datatype, info->algorithm, info->protocol);
|
||||
}
|
||||
|
||||
work->connIndex = 0;
|
||||
|
||||
@@ -19,9 +19,8 @@ struct ncclDevRedOpFull {
|
||||
uint64_t scalarArg;
|
||||
};
|
||||
|
||||
#define FUNC_INDEX_P2P (ncclNumTypes+NCCL_NUM_FUNCTIONS*NCCL_NUM_ALGORITHMS*NCCL_NUM_PROTOCOLS*ncclNumTypes*ncclNumDevRedOps)
|
||||
#define FUNC_INDEX_ALLTOALL_PIVOT (FUNC_INDEX_P2P+1)
|
||||
#define FUNC_INDEX(func, devredop, ncclType, al, pr) ((((((func)*ncclNumDevRedOps + (devredop))*ncclNumTypes) + (ncclType))*NCCL_NUM_ALGORITHMS+(al))*NCCL_NUM_PROTOCOLS+(pr))
|
||||
#define FUNC_INDEX_P2P 1015
|
||||
#define FUNC_INDEX_ALLTOALL_PIVOT 675
|
||||
|
||||
#define NCCL_FUNC_NAME(func, algo, proto, devredop, type) \
|
||||
ncclFunction_##func##_##algo##_##proto##_##devredop##_##type
|
||||
@@ -38,79 +37,13 @@ struct ncclDevRedOpFull {
|
||||
#define NCCL_IMPL_NAME(func, algo, proto) \
|
||||
nccl##func##algo##proto
|
||||
|
||||
/* Declare all collective operations */
|
||||
#if defined(USE_INDIRECT_FUNCTION_CALL) && !defined(__gfx940__) && !defined(__gfx941__) && !defined(__gfx942__)
|
||||
#define DECL5(func, algo, proto, devredop, type) \
|
||||
extern __device__ void NCCL_FUNC_NAME(func, algo, proto, devredop, type)(); \
|
||||
extern __global__ void NCCL_KERN_NAME(func, algo, proto, devredop, type)(struct ncclDevComm* comm, uint64_t channelMask, struct ncclWork* workHead); \
|
||||
extern __global__ void NCCL_KERN_NAME_DEBUG(func, algo, proto, devredop, type)(struct ncclDevComm* comm, uint64_t channelMask, struct ncclWork* workHead);
|
||||
#else
|
||||
#define DECL5(func, algo, proto, devredop, type) \
|
||||
extern __device__ __attribute__((noinline)) void NCCL_FUNC_NAME(func, algo, proto, devredop, type)(); \
|
||||
extern __global__ void NCCL_KERN_NAME(func, algo, proto, devredop, type)(struct ncclDevComm* comm, uint64_t channelMask, struct ncclWork* workHead); \
|
||||
extern __global__ void NCCL_KERN_NAME_DEBUG(func, algo, proto, devredop, type)(struct ncclDevComm* comm, uint64_t channelMask, struct ncclWork* workHead);
|
||||
// Declare rccl main/general kernel
|
||||
extern __global__ void rccl_main_kernel(struct ncclDevComm* comm, uint64_t channelMask, struct ncclWork* workHead);
|
||||
#ifdef ENABLE_COLLTRACE
|
||||
extern __global__ void rccl_main_kernel_debug(struct ncclDevComm* comm, uint64_t channelMask, struct ncclWork* workHead);
|
||||
#endif
|
||||
|
||||
#define SINGLE_ARG(...) __VA_ARGS__
|
||||
#define CONCAT(a,b) a##b
|
||||
#define MACRO_IF(cond, t, f) CONCAT(MACRO_IF_, cond)(SINGLE_ARG(t), SINGLE_ARG(f))
|
||||
#define MACRO_IF_0(t, f) f
|
||||
#define MACRO_IF_1(t, f) t
|
||||
|
||||
#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, devredop, type, undef) \
|
||||
DECL4(func, RING, devredop, type, undef) \
|
||||
DECL4(func, TREE, devredop, type, undef) \
|
||||
DECL4(func, COLLNET_DIRECT, devredop, type, undef) \
|
||||
DECL4(func, COLLNET_CHAIN, devredop, type, undef) \
|
||||
DECL4(func, NVLS, devredop, type, undef) \
|
||||
DECL4(func, NVLS_TREE, devredop, type, undef)
|
||||
|
||||
#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, 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, /*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)
|
||||
|
||||
DECL2(Broadcast, Sum, /*undefForFloat=*/0)
|
||||
DECL(Reduce)
|
||||
DECL2(AllGather, Sum, /*undefForFloat=*/0)
|
||||
DECL(ReduceScatter)
|
||||
DECL(AllReduce)
|
||||
DECL5(SendRecv, RING, SIMPLE, Sum, int8_t)
|
||||
DECL5(AllToAllPivot, RING, SIMPLE, Sum, int8_t)
|
||||
|
||||
// Declare OneRankReduce
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, int8_t)();
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, uint8_t)();
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, int32_t)();
|
||||
@@ -118,9 +51,7 @@ extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, uint32_t)();
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, int64_t)();
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, uint64_t)();
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, half)();
|
||||
#if defined(RCCL_BFLOAT16)
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, rccl_bfloat16)();
|
||||
#endif
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, float)();
|
||||
extern __device__ void NCCL_ONERANK_REDUCE_NAME(PreMulSum, double)();
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@
|
||||
#include "nccl.h"
|
||||
#include "rccl_bfloat16.h"
|
||||
#include "align.h"
|
||||
#include "collectives.h"
|
||||
#if defined(ENABLE_NPKIT)
|
||||
#include "npkit/npkit_struct.h"
|
||||
#endif
|
||||
@@ -483,4 +484,62 @@ __host__ __device__ constexpr int ncclShmemDynamicSize(int cudaArch = NCCL_CUDA_
|
||||
return cudaArch < 700 ? 0 : ncclShmemScratchWarpSize(cudaArch)*(NCCL_MAX_NTHREADS/WARP_SIZE);
|
||||
}
|
||||
|
||||
// Map the rowIdx to funcIdx
|
||||
extern int const ncclFuncRowToId[];
|
||||
|
||||
// `ncclFuncIndex()` needs to be in sync with 'ALL_COLLS' in Generate.cmake
|
||||
inline int ncclFuncId(int coll, int devRedOp, int type, int algo, int proto) {
|
||||
int row = 0;
|
||||
|
||||
// RING / <all_protos> / Sum / int8_t
|
||||
if (coll == ncclFuncAllGather) {
|
||||
row += proto;
|
||||
goto have_row;
|
||||
}
|
||||
row += NCCL_NUM_PROTOCOLS;
|
||||
|
||||
// <all_algos> / <all_protos> / <all_redops> / <all_types>
|
||||
if (coll == ncclFuncAllReduce) {
|
||||
row += (((algo * NCCL_NUM_PROTOCOLS + proto) * ncclNumDevRedOps + devRedOp) * ncclNumTypes + type) - /*floats for each SumPostDiv*/ 4 * (algo * NCCL_NUM_PROTOCOLS + proto);
|
||||
goto have_row;
|
||||
}
|
||||
row += (NCCL_NUM_ALGORITHMS - 2) * NCCL_NUM_PROTOCOLS * (ncclNumDevRedOps * ncclNumTypes - /*floats for each SumPostDiv*/ 4);
|
||||
|
||||
// RING / SIMPLE / Sum / int8_t
|
||||
if (coll == ncclFuncAllToAllPivot) goto have_row;
|
||||
row += 1;
|
||||
|
||||
// RING / <all_protos> / Sum / int8_t
|
||||
if (coll == ncclFuncBroadcast) {
|
||||
row += proto;
|
||||
goto have_row;
|
||||
}
|
||||
row += NCCL_NUM_PROTOCOLS;
|
||||
|
||||
// RING / <all_protos> / <all_redops> / <all_types>
|
||||
if (coll == ncclFuncReduce) {
|
||||
row += ((proto * ncclNumDevRedOps + devRedOp) * ncclNumTypes + type) - /*floats for each SumPostDiv*/ 4 * proto;
|
||||
goto have_row;
|
||||
}
|
||||
row += NCCL_NUM_PROTOCOLS * (ncclNumDevRedOps * ncclNumTypes - /*floats for each SumPostDiv*/ 4);
|
||||
|
||||
// RING / <all_protos> / <all_redops> / <all_types>
|
||||
if (coll == ncclFuncReduceScatter) {
|
||||
row += ((proto * ncclNumDevRedOps + devRedOp) * ncclNumTypes + type) - /*floats for each SumPostDiv*/ 4 * proto;
|
||||
goto have_row;
|
||||
}
|
||||
row += NCCL_NUM_PROTOCOLS * (ncclNumDevRedOps * ncclNumTypes - /*floats for each SumPostDiv*/ 4);
|
||||
|
||||
// RING / SIMPLE / Sum / int8_t
|
||||
if (coll == ncclFuncSendRecv) goto have_row;
|
||||
row += 1;
|
||||
|
||||
have_row:
|
||||
return ncclFuncRowToId[row];
|
||||
}
|
||||
|
||||
inline int ncclFuncId_P2p() { return ncclFuncRowToId[FUNC_INDEX_P2P]; }
|
||||
|
||||
inline int ncclFuncId_AllToAllPivot() { return ncclFuncRowToId[FUNC_INDEX_ALLTOALL_PIVOT]; }
|
||||
|
||||
#endif
|
||||
|
||||
@@ -10,6 +10,7 @@
|
||||
#include "comm.h"
|
||||
#include "group.h"
|
||||
#include "collectives.h"
|
||||
#include "common.h"
|
||||
#include "utils.h"
|
||||
|
||||
#define NCCL_MIN_CHANNEL_SIZE (NCCL_LL_THREAD_THRESHOLD*64)
|
||||
|
||||
+58
-20
@@ -18,6 +18,7 @@
|
||||
#include "enqueue.h"
|
||||
#include "graph.h"
|
||||
#include "argcheck.h"
|
||||
#include "devcomm.h"
|
||||
#if defined(ENABLE_NPKIT)
|
||||
#include "npkit/npkit.h"
|
||||
#endif
|
||||
@@ -31,6 +32,7 @@
|
||||
#include <sys/types.h>
|
||||
#include <sys/stat.h>
|
||||
#include <unistd.h>
|
||||
#include <cstdarg>
|
||||
#include "graph/topo.h"
|
||||
#include "graph/xml.h"
|
||||
#include "archinfo.h"
|
||||
@@ -54,7 +56,7 @@
|
||||
#define NCCL_GROUP_CUDA_STREAM 1 // CGMD: CUDA 9.0,9.1 Need to use an internal CUDA stream
|
||||
#endif
|
||||
|
||||
const char* ncclFuncStr[NCCL_NUM_FUNCTIONS+2] = { "Broadcast", "Reduce", "AllGather", "ReduceScatter", "AllReduce", "SendRecv", "AllToAllPivot" };
|
||||
const char* ncclFuncStr[NCCL_NUM_FUNCTIONS+2] = { "AllGather", "AllReduce", "AllToAllPivot", "Broadcast", "Reduce", "ReduceScatter", "SendRecv"};
|
||||
const char* ncclAlgoStr[NCCL_NUM_ALGORITHMS] = { "Tree", "Ring", "CollNetDirect", "CollNetChain", "NVLS", "NVLSTree" };
|
||||
const char* ncclProtoStr[NCCL_NUM_PROTOCOLS] = { "LL", "LL128", "Simple" };
|
||||
const char* ncclDevRedOpStr[ncclNumDevRedOps] = { "Sum", "Prod", "Max", "Min", "PreMulSum", "SumPostDiv" };
|
||||
@@ -177,6 +179,18 @@ void NCCL_NO_OPTIMIZE commPoison(ncclComm_t comm) {
|
||||
RCCL_PARAM(KernelCollTraceEnable, "KERNEL_COLL_TRACE_ENABLE", 0);
|
||||
|
||||
#ifdef ENABLE_COLLTRACE
|
||||
#define MAX_NAME_LENGTH 64
|
||||
// Helper function to generate function names and update funcIdx
|
||||
void generateFunctionName(char* func_names, int& funcIdx, const char* format, ...) {
|
||||
char* line = func_names + MAX_NAME_LENGTH * funcIdx;
|
||||
va_list args;
|
||||
va_start(args, format);
|
||||
vsnprintf(line, MAX_NAME_LENGTH, format, args);
|
||||
va_end(args);
|
||||
funcIdx++;
|
||||
}
|
||||
|
||||
// Should be in sync with 'ALL_COLLS' in Generator.cmake
|
||||
void *ncclCommThreadMain(void *arg) {
|
||||
ncclComm_t comm = (ncclComm_t)arg;
|
||||
int head[MAXCHANNELS];
|
||||
@@ -184,29 +198,53 @@ void *ncclCommThreadMain(void *arg) {
|
||||
|
||||
memset(head, 0, sizeof(int)*MAXCHANNELS);
|
||||
vega_gpu_rtc_freq = GetDeviceWallClockRateInKhz(comm->cudaDev) * 1.0E3;
|
||||
#define MAX_NAME_LENGTH 64
|
||||
char* func_names = (char *)malloc(MAX_NAME_LENGTH*(FUNC_INDEX_P2P+2));
|
||||
for (int func = 0; func < NCCL_NUM_FUNCTIONS; func++) {
|
||||
for (int al = 0; al < NCCL_NUM_ALGORITHMS; al++) {
|
||||
for (int type = 0; type < ncclNumTypes; type++) {
|
||||
for (int pr = 0; pr < NCCL_NUM_PROTOCOLS; pr++) {
|
||||
for (int devredop = 0; devredop < ncclNumDevRedOps; devredop++) {
|
||||
char* line = func_names+MAX_NAME_LENGTH*FUNC_INDEX(func, devredop, type, al, pr);
|
||||
sprintf(line, "%s%s%s%s%s", ncclFuncStr[func], ncclAlgoStr[al], ncclProtoStr[pr],
|
||||
ncclDevRedOpStr[devredop], ncclTypeStr[type]);
|
||||
}
|
||||
char* func_names = (char *)malloc(MAX_NAME_LENGTH*(ncclFuncId_P2p()+/*OneRankReduce*/11));
|
||||
int funcIdx = 0;
|
||||
// AllGather --> RING / <all_protos> / Sum / int8_t
|
||||
for (int pr = 0; pr < NCCL_NUM_PROTOCOLS; pr++) {
|
||||
generateFunctionName(func_names, funcIdx, "AllGatherRing%sSum_i8", ncclProtoStr[pr]);
|
||||
}
|
||||
// AllReduce --> <all_algos> / <all_protos> / <all_redops> / <all_types>
|
||||
for (int al = 0; al < NCCL_NUM_ALGORITHMS - 2; al++) {
|
||||
for (int pr = 0; pr < NCCL_NUM_PROTOCOLS; pr++) {
|
||||
for (int redop = 0; redop < ncclNumDevRedOps; redop++) {
|
||||
for (int ty = 0; ty < ncclNumTypes; ty++) {
|
||||
if (redop == 5 && ty > 5) continue;
|
||||
generateFunctionName(func_names, funcIdx, "AllReduce%s%s%s%s", ncclAlgoStr[al], ncclProtoStr[pr], ncclDevRedOpStr[redop], ncclTypeStr[ty]);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
for (int type = 0; type < ncclNumTypes; type++) {
|
||||
char* line = func_names+MAX_NAME_LENGTH*(FUNC_INDEX_P2P-ncclNumTypes+type);
|
||||
sprintf(line, "OneRankReducePreMulSum%s", ncclTypeStr[type]);
|
||||
// AllToAllPivot --> RING / SIMPLE / Sum / int8_t
|
||||
generateFunctionName(func_names, funcIdx, "AllToAllPivotRingSimpleSum_i8");
|
||||
// Broadcast --> RING / <all_protos> / Sum / int8_t
|
||||
for (int pr = 0; pr < NCCL_NUM_PROTOCOLS; pr++) {
|
||||
generateFunctionName(func_names, funcIdx, "BroadcastRing%sSum_i8", ncclProtoStr[pr]);
|
||||
}
|
||||
// Reduce --> RING / <all_protos> / <all_redops> / <all_types>
|
||||
for (int pr = 0; pr < NCCL_NUM_PROTOCOLS; pr++) {
|
||||
for (int redop = 0; redop < ncclNumDevRedOps; redop++) {
|
||||
for (int ty = 0; ty < ncclNumTypes; ty++) {
|
||||
if (redop == 5 && ty > 5) continue;
|
||||
generateFunctionName(func_names, funcIdx, "ReduceRing%s%s%s", ncclProtoStr[pr], ncclDevRedOpStr[redop], ncclTypeStr[ty]);
|
||||
}
|
||||
}
|
||||
}
|
||||
// ReduceScatter --> RING / <all_protos> / <all_redops> / <all_types>
|
||||
for (int pr = 0; pr < NCCL_NUM_PROTOCOLS; pr++) {
|
||||
for (int redop = 0; redop < ncclNumDevRedOps; redop++) {
|
||||
for (int ty = 0; ty < ncclNumTypes; ty++) {
|
||||
if (redop == 5 && ty > 5) continue;
|
||||
generateFunctionName(func_names, funcIdx, "ReduceScatterRing%s%s%s", ncclProtoStr[pr], ncclDevRedOpStr[redop], ncclTypeStr[ty]);
|
||||
}
|
||||
}
|
||||
}
|
||||
// SendRecv --> RING / SIMPLE / Sum / int8_t
|
||||
generateFunctionName(func_names, funcIdx, "SendRecvRingSimpleSum_i8");
|
||||
// OneRankReduce --> PreMulSum / <all_types>
|
||||
for (int ty = 0; ty < ncclNumTypes; ty++) {
|
||||
generateFunctionName(func_names, funcIdx, "OneRankReducePreMulSum%s", ncclTypeStr[ty]);
|
||||
}
|
||||
char* line = func_names+MAX_NAME_LENGTH*FUNC_INDEX_P2P;
|
||||
sprintf(line, "SendRecvRingSimpleSum_i8");
|
||||
line += MAX_NAME_LENGTH;
|
||||
sprintf(line, "AllToAllPivotRingSimpleSum_i8");
|
||||
do {
|
||||
for (int channel = 0; channel < MAXCHANNELS; channel++) {
|
||||
int tail = comm->collTraceTail[channel].tail%COLLTRACE_NUM_ITEMS;
|
||||
@@ -232,7 +270,7 @@ void *ncclCommThreadMain(void *arg) {
|
||||
(double)(td->timeStamp)/vega_gpu_rtc_freq, comm->rank, td->bid,
|
||||
fIdx, td->data_0, td->opCount, td->data_1);
|
||||
} else {
|
||||
if (fIdx == FUNC_INDEX_P2P || type == ncclCollTraceP2pElemType)
|
||||
if (fIdx == ncclFuncId_P2p() || type == ncclCollTraceP2pElemType)
|
||||
sprintf(line, "## [%012.6f] [%02d:%02d] %06x-%06x", (double)(td->timeStamp)/vega_gpu_rtc_freq, comm->rank, td->bid, td->p2pOpCount[0], td->p2pOpCount[1]);
|
||||
else
|
||||
sprintf(line, "## [%012.6f] [%02d:%02d] %06lx", (double)(td->timeStamp)/vega_gpu_rtc_freq, comm->rank, td->bid, td->opCount);
|
||||
|
||||
Reference in New Issue
Block a user