Revert "Port alltoall[v]" (#325)

This reverts commit 2c49121171.

[ROCm/rccl commit: 8e180cf087]
Cette révision appartient à :
Wenkai Du
2021-03-06 13:59:31 -08:00
révisé par GitHub
Parent bcf4ecb0e3
révision b7253710ca
25 fichiers modifiés avec 65 ajouts et 451 suppressions
+8 -53
Voir le fichier
@@ -280,13 +280,6 @@ static ncclResult_t getAlgoInfo(struct ncclInfo* info) {
struct ncclComm* comm = info->comm;
float minTime = 3600000000.0; // Hopefully no operation will take an hour to complete.
// Find algorithm / protocol.
if (info->coll == ncclFuncAllToAll || info->coll == ncclFuncAllToAllv) {
info->algorithm = NCCL_ALGO_RING;
info->protocol = NCCL_PROTO_SIMPLE;
info->nChannels = comm->nChannels;
info->nThreads = NCCL_MAX_NTHREADS;
return ncclSuccess;
}
info->algorithm = -1;
info->protocol = -1;
int nAlgos = NCCL_NUM_ALGORITHMS;
@@ -348,9 +341,6 @@ static ncclResult_t getPatternInfo(struct ncclInfo* info) {
info->pattern = ncclPatternRing; break;
case ncclFuncAllReduce:
info->pattern = info->algorithm == NCCL_ALGO_COLLNET ? ncclPatternCollTreeUp : info->algorithm == NCCL_ALGO_TREE ? ncclPatternTreeUpDown : ncclPatternRingTwice; break;
case ncclFuncAllToAll:
case ncclFuncAllToAllv:
info->pattern = ncclPatternAll; break;
default:
WARN("Unknown pattern for collective %d algorithm %d", info->coll, info->algorithm);
return ncclInternalError;
@@ -372,9 +362,6 @@ static ncclResult_t getLoopInfo(struct ncclInfo* info) {
info->nstepsPerLoop = info->comm->nRanks-1; info->nchunksPerLoop = info->comm->nRanks; break;
case ncclPatternRingTwice:
info->nstepsPerLoop = 2*(info->comm->nRanks-1); info->nchunksPerLoop = info->comm->nRanks; break;
case ncclPatternAll:
info->nstepsPerLoop = 1;
info->nchunksPerLoop = info->comm->nRanks; break;
default:
WARN("Unknown pattern %d", info->pattern);
return ncclInternalError;
@@ -390,21 +377,12 @@ static ncclResult_t computeColl(struct ncclInfo* info /* input */, struct ncclWo
NCCLCHECK(getPatternInfo(info));
NCCLCHECK(getLoopInfo(info));
if ((info->coll == ncclFuncAllToAll || info->coll == ncclFuncAllToAllv)
&& info->comm->topo->nodes[GPU].count == info->comm->topo->nRanks && (info->comm->topo->type & RCCL_TOPO_4P2H_ROME))
info->nChannels = info->comm->p2pnChannels;
work->opCount = info->comm->opCount;
work->sendbuff = info->sendbuff;
work->recvbuff = info->recvbuff;
if (info->coll == ncclFuncAllToAllv) {
work->a2av.count = info->count;
work->a2av.nChannels = info->nChannels;
} else {
work->coll.root = info->root;
work->coll.count = info->count;
work->coll.nChannels = info->nChannels;
}
work->coll.root = info->root;
work->coll.count = info->count;
work->coll.nChannels = info->nChannels;
work->nThreads = info->nThreads;
work->funcIndex = FUNC_INDEX(info->coll, info->op, info->datatype, info->algorithm, info->protocol);
@@ -522,7 +500,7 @@ ncclResult_t ncclSaveKernel(struct ncclInfo* info) {
info->comm->myParams->blockDim.x = std::max<unsigned>(info->comm->myParams->blockDim.x, info->nThreads);
int nChannels = (info->coll == ncclFuncAllToAllv) ? work.a2av.nChannels : work.coll.nChannels;
int nChannels = work.coll.nChannels;
int nSubChannels = (info->pattern == ncclPatternCollTreeUp || info->pattern == ncclPatternCollTreeDown) ? 2 : 1;
for (int bid=0; bid<nChannels*nSubChannels; bid++) {
@@ -539,11 +517,6 @@ ncclResult_t ncclSaveKernel(struct ncclInfo* info) {
if (proxyArgs.nsteps) NCCLCHECK(ncclProxySaveColl(&proxyArgs, info->pattern, info->root, info->comm->nRanks));
info->comm->myParams->gridDim.x++;
if (info->coll == ncclFuncAllToAllv) {
work.a2av.bid = bid % work.a2av.nChannels;
} else {
work.coll.bid = bid % nChannels;
}
// [RCCL] Setup pointers to where all the input/output pointers will be
if (info->protocol == NCCL_PROTO_CLIQUE) {
@@ -551,16 +524,8 @@ ncclResult_t ncclSaveKernel(struct ncclInfo* info) {
}
// [/RCCL]
struct ncclWork* w;
NCCLCHECK(getNextOp(channel, &w, &work));
if (info->coll == ncclFuncAllToAllv) {
struct ncclWorkElem* e = w->elems;
size_t* params = channel->a2avParams + info->comm->nRanks*4*e->index;
memcpy(params, info->sendcounts, sizeof(size_t*)*(info->comm->nRanks));
memcpy(params+info->comm->nRanks, info->sdispls, sizeof(size_t*)*(info->comm->nRanks));
memcpy(params+info->comm->nRanks*2, info->recvcounts, sizeof(size_t*)*(info->comm->nRanks));
memcpy(params+info->comm->nRanks*3, info->rdispls, sizeof(size_t*)*(info->comm->nRanks));
}
work.coll.bid = bid % nChannels;
NCCLCHECK(getNextOp(channel, NULL, &work));
}
info->comm->opCount++;
return ncclSuccess;
@@ -714,12 +679,7 @@ ncclResult_t ncclEnqueueCheck(struct ncclInfo* info) {
NCCLCHECKGOTO(ncclAsyncColl(info->comm), ret, end);
NCCLCHECKGOTO(checkSetStream(info), ret, end);
if (info->coll == ncclFuncAllToAllv)
INFO(NCCL_COLL,"%s: opCount %lx sendbuff %p sendcounts %p sdispls %p recvbuff %p recvcounts %p rdispls %p datatype %d typesize %zi op %d root %d comm %p [nranks=%d] stream %p",
info->opName, info->comm->opCount, info->sendbuff, info->sendcounts, info->sdispls, info->recvbuff, info->recvcounts, info->rdispls,
info->datatype, info->count, info->op, info->root, info->comm, info->comm->nRanks, info->stream);
else if (info->coll != ncclFuncSendRecv)
INFO(NCCL_COLL,"%s: opCount %lx sendbuff %p recvbuff %p count %zi datatype %d op %d root %d comm %p [nranks=%d] stream %p",
INFO(NCCL_COLL,"%s: opCount %lx sendbuff %p recvbuff %p count %zi datatype %d op %d root %d comm %p [nranks=%d] stream %p",
info->opName, info->comm->opCount, info->sendbuff, info->recvbuff, info->count,
info->datatype, info->op, info->root, info->comm, info->comm->nRanks, info->stream);
@@ -741,12 +701,7 @@ end:
NCCLCHECK(ArgsCheck(info));
NCCLCHECK(checkSetStream(info));
if (info->coll == ncclFuncAllToAllv)
INFO(NCCL_COLL,"%s: opCount %lx sendbuff %p sendcounts %p sdispls %p recvbuff %p recvcounts %p rdispls %p datatype %d typesize %zi op %d root %d comm %p [nranks=%d] stream %p",
info->opName, info->comm->opCount, info->sendbuff, info->sendcounts, info->sdispls, info->recvbuff, info->recvcounts, info->rdispls,
info->datatype, info->count, info->op, info->root, info->comm, info->comm->nRanks, info->stream);
else
INFO(NCCL_COLL,"%s: opCount %lx sendbuff %p recvbuff %p count %zi datatype %d op %d root %d comm %p [nranks=%d] stream %p",
INFO(NCCL_COLL,"%s: opCount %lx sendbuff %p recvbuff %p count %zi datatype %d op %d root %d comm %p [nranks=%d] stream %p",
info->opName, info->comm->opCount, info->sendbuff, info->recvbuff, info->count,
info->datatype, info->op, info->root, info->comm, info->comm->nRanks, info->stream);