recreated pr 914 to work with current develop branch (#979)

This commit is contained in:
akolliasAMD
2023-11-28 16:33:47 -07:00
committed by GitHub
parent c71bae1608
commit 56ce9ef05f
5 changed files with 42 additions and 7 deletions
+3 -2
View File
@@ -215,10 +215,11 @@ static ncclResult_t mscclInternalSchedulerSelectAlgo(struct mscclSchedulerParam*
mscclStatus& status = mscclGetStatus();
param->scheduled = false;
// Current MSCCL doesn't support pre/post op
/*// Current MSCCL doesn't support pre/post op
if (param->op >= ncclAvg) {
return ncclSuccess;
}
}*/
// Whether the algorithm is in-place
bool isInPlace = false;
+16 -2
View File
@@ -309,6 +309,18 @@ static ncclResult_t hostToDevRedOp(
MSCCL_KERNEL_ENTRY_DEVREDOP_TYPE(devredop, double, fullOps), \
MSCCL_KERNEL_ENTRY_DEVREDOP_TYPE(devredop, rccl_bfloat16, fullOps)
#define MSCCL_KERNEL_ENTRY_DEVREDOP_NOFLOAT(devredop, fullOps) \
MSCCL_KERNEL_ENTRY_DEVREDOP_TYPE(devredop, int8_t, fullOps), \
MSCCL_KERNEL_ENTRY_DEVREDOP_TYPE(devredop, uint8_t, fullOps), \
MSCCL_KERNEL_ENTRY_DEVREDOP_TYPE(devredop, int32_t, fullOps), \
MSCCL_KERNEL_ENTRY_DEVREDOP_TYPE(devredop, uint32_t, fullOps), \
MSCCL_KERNEL_ENTRY_DEVREDOP_TYPE(devredop, int64_t, fullOps), \
MSCCL_KERNEL_ENTRY_DEVREDOP_TYPE(devredop, uint64_t, fullOps), \
MSCCL_KERNEL_ENTRY_DEVREDOP_NULL(), \
MSCCL_KERNEL_ENTRY_DEVREDOP_NULL(), \
MSCCL_KERNEL_ENTRY_DEVREDOP_NULL(), \
MSCCL_KERNEL_ENTRY_DEVREDOP_NULL()
#define MSCCL_KERNEL_ENTRY() \
MSCCL_KERNEL_ENTRY_DEVREDOP(Sum, false), \
MSCCL_KERNEL_ENTRY_DEVREDOP(Prod, false), \
@@ -317,10 +329,12 @@ static ncclResult_t hostToDevRedOp(
MSCCL_KERNEL_ENTRY_DEVREDOP(Sum, true), \
MSCCL_KERNEL_ENTRY_DEVREDOP(Prod, true), \
MSCCL_KERNEL_ENTRY_DEVREDOP(Max, true), \
MSCCL_KERNEL_ENTRY_DEVREDOP(Min, true)
MSCCL_KERNEL_ENTRY_DEVREDOP(Min, true), \
MSCCL_KERNEL_ENTRY_DEVREDOP(PreMulSum, true), \
MSCCL_KERNEL_ENTRY_DEVREDOP_NOFLOAT(SumPostDiv, true)
// Except for ncclDevPreMulSum and ncclDevSumPostDiv required by ncclAvg
void* mscclKernelEntries[(ncclNumDevRedOps - 2) * ncclNumTypes * NCCL_NUM_PROTOCOLS * 2] = {
void* mscclKernelEntries[ncclNumDevRedOps * ncclNumTypes * NCCL_NUM_PROTOCOLS * 2] = {
#ifdef COMPILE_MSCCL_KERNEL
MSCCL_KERNEL_ENTRY()
#endif