/*************************************************************************
 * Copyright (c) 2016-2022, NVIDIA CORPORATION. All rights reserved.
 *
 * See LICENSE.txt for license information
 ************************************************************************/

#include "cuda_runtime.h"
#include "common.h"
#if NCCL_VERSION_CODE >= NCCL_VERSION(2,28,0)
#include "nccl_device.h"
#endif

#if defined(NCCL_OS_LINUX)
#pragma weak ncclAlltoAll
#endif

void AlltoAllGetCollByteCount(size_t *sendcount, size_t *recvcount, size_t *paramcount, size_t *sendInplaceOffset, size_t *recvInplaceOffset, size_t count, size_t eltSize, int nranks) {
  *paramcount = (count/nranks) & ~(16/eltSize - 1);
  *sendcount = nranks*(*paramcount);
  *recvcount = *sendcount;
  *sendInplaceOffset = 0;
  *recvInplaceOffset = 0;
}

testResult_t AlltoAllInitData(struct threadArgs* args, ncclDataType_t type, ncclRedOp_t op, int root, int rep, int in_place) {
  size_t sendcount = args->sendBytes / wordSize(type);
  size_t recvcount = args->expectedBytes / wordSize(type);
  int nranks = args->nProcs*args->nThreads*args->nGpus;

  for (int i=0; i<args->nGpus; i++) {
    CUDACHECK(cudaSetDevice(args->gpus[i]));
    int rank = ((args->proc*args->nThreads + args->thread)*args->nGpus + i);
    CUDACHECK(cudaMemset(args->recvbuffs[i], 0, args->expectedBytes));
    void* data = in_place ? args->recvbuffs[i] : args->sendbuffs[i];
    TESTCHECK(InitData(data, sendcount, 0, type, ncclSum, 33*rep + rank, 1, 0));
    for (int j=0; j<nranks; j++) {
      size_t partcount = sendcount/nranks;
      TESTCHECK(InitData((char*)args->expected[i] + j*partcount*wordSize(type), partcount, rank*partcount, type, ncclSum, 33*rep + j, 1, 0));
    }
    CUDACHECK(cudaDeviceSynchronize());
  }
  // We don't support in-place alltoall
  args->reportErrors = in_place ? 0 : 1;
  return testSuccess;
}

void AlltoAllGetBw(size_t count, size_t typesize, double sec, double* algBw, double* busBw, int nranks) {
  double baseBw = (double)(count * nranks * typesize) / 1.0E9 / sec;

  *algBw = baseBw;
  double factor = ((double)(nranks-1))/((double)(nranks));
  *busBw = baseBw * factor;
}

#if NCCL_VERSION_CODE >= NCCL_VERSION(2,29,0)
// set devComm reqs for alltoall device kernels
testResult_t AlltoAllGetDevCommRequirements(int deviceImpl, ncclDevCommRequirements* reqs, ncclComm_t comm) {
  if (!reqs || !comm) return testInternalError;

  ncclCommProperties_t commProperties = NCCL_COMM_PROPERTIES_INITIALIZER;
  if (ncclCommQueryProperties(comm, &commProperties) != ncclSuccess) {
    return testNcclError;
  }

  switch(deviceImpl) {
    case 1: // NvlAlltoAllKernel
    case 2: // NvlAlltoAllKernelOptimized
      if (commProperties.nRanks != ncclTeamLsa(comm).nRanks) {
        fprintf(stderr, "DeviceImplementation 1 and 2 requires CUDA P2P "
                        "connectivity across all ranks. Not all ranks of this "
                        "communicator have P2P connectivity.\n");
        return testInvalidUsage;
      }
      reqs->lsaBarrierCount = deviceCtaCount;
      return testSuccess;
    #if defined(NCCL_OS_LINUX)
    case 5: // GinAlltoAllKernelMultiContext
      reqs->ginContextCount = deviceCtaCount;           // 1 ctx per CTA
      // fall through
    case 3: // GinAlltoAllKernel
      if (deviceImpl == 3) reqs->ginContextCount = 1;   // 1 ctx per GIN connection
#if NCCL_VERSION_CODE >= NCCL_VERSION(2, 30, 0)
      reqs->worldGinBarrierCount = deviceCtaCount;
#endif
      // fall through
    case 4: // HybridAlltoAllKernel (LSA+GIN)
      if (deviceImpl == 4) reqs->ginContextCount = 1;   // 1 ctx per GIN connection
      if (commProperties.ginType == NCCL_GIN_TYPE_NONE) {
        fprintf(stderr, "This test requires GIN support, but GIN support is not enabled for this communicator.\n");
        return testInvalidUsage;
      }
#if NCCL_VERSION_CODE >= NCCL_VERSION(2, 30, 0)
      if (deviceImpl == 4)
#endif
        reqs->barrierCount = deviceCtaCount;
      reqs->ginSignalCount = deviceCtaCount;
#if NCCL_VERSION_CODE >= NCCL_VERSION(2, 30, 7)
      reqs->ginStrongSignalsRequired = false;
      reqs->ginVaSignalsRequired = false;
#endif
#if NCCL_VERSION_CODE >= NCCL_VERSION(2, 29, 7)
      reqs->ginConnectionType = NCCL_GIN_CONNECTION_FULL;
#else
      reqs->ginForceEnable = true;
#endif
      return testSuccess;
    #endif
    default:
      return testNotImplemented;
  }
}
#elif NCCL_VERSION_CODE >= NCCL_VERSION(2,28,0)
// set devComm reqs for alltoall device kernels
bool AlltoAllGetDevCommRequirements(int deviceImpl, ncclDevCommRequirements* reqs) {
  if (!reqs) return false;
  memset(reqs, 0, sizeof(*reqs));

  switch(deviceImpl) {
    case 1: // NvlAlltoAllKernel
    case 2: // NvlAlltoAllKernelOptimized
      reqs->lsaBarrierCount = deviceCtaCount;
      return true;
#if NCCL_VERSION_CODE >= NCCL_VERSION(2,28,7) && defined(NCCL_OS_LINUX)
    case 5: // GinAlltoAllKernelMultiContext
      reqs->ginContextCount = deviceCtaCount;   // 1 ctx per CTA
      // fall through
    case 3: // GinAlltoAllKernel
    case 4: // HybridAlltoAllKernel (LSA+GIN)
      if (deviceImpl == 3 || deviceImpl == 4) reqs->ginContextCount = 1;   // 1 ctx per GIN connection
      reqs->barrierCount = deviceCtaCount;
      reqs->ginSignalCount = deviceCtaCount;
      return true;
#endif
    default:
      return false;
  }
}
#endif

#if NCCL_VERSION_CODE >= NCCL_VERSION(2,28,0)
// shared scalar AlltoAll implementation used by both kernels
template <typename T>
__device__ void AlltoAllScalarImpl(ncclWindow_t sendwin, size_t sendoffset, ncclWindow_t recvwin, size_t recvoffset, size_t count, int rank, int nRanks, int tid, int nthreads) {
  T* sendPtr = (T*)ncclGetLsaPointer(sendwin, sendoffset, rank);

  for (size_t offset = tid; offset < count; offset += nthreads) {
    for (int peer = 0; peer < nRanks; peer++) {
      T value = sendPtr[peer * count + offset];
      T* recvPtr = (T*)ncclGetLsaPointer(recvwin, recvoffset, peer);
      recvPtr[rank * count + offset] = value;
    }
  }
}

// Device implementation #1 - simple NVL kernel
template <typename T>
__global__ void NvlAlltoAllKernel(ncclWindow_t sendwin, size_t sendoffset, ncclWindow_t recvwin, size_t recvoffset, size_t count, int root, struct ncclDevComm devComm) {
  ncclLsaBarrierSession<ncclCoopCta> bar { ncclCoopCta(), devComm, ncclTeamLsa(devComm), devComm.lsaBarrier, blockIdx.x };
  bar.sync(ncclCoopCta(), cuda::memory_order_acquire);

  int rank = devComm.rank, nRanks = devComm.nRanks;
  int tid = threadIdx.x + blockDim.x * blockIdx.x;
  int nthreads = blockDim.x * gridDim.x;

  AlltoAllScalarImpl<T>(sendwin, sendoffset, recvwin, recvoffset, count, rank, nRanks, tid, nthreads);

  bar.sync(ncclCoopCta(), cuda::memory_order_release);
}

// shared across -D 2 and -D 4 kernels - optimized NVL kernel using vectorization and unrolling
template <typename T>
__device__ void AlltoAllLsaVecImpl(ncclWindow_t sendwin, size_t sendoffset, ncclWindow_t recvwin, size_t recvoffset,
    size_t count, int worldRank, int startLsa, int lsaSize, int tid, int nthreads) {
  using TN = uint4; // Alltoall is type insensitive, so using generic uint4 for data transport
  constexpr int VECTOR_FACTOR = sizeof(TN) / sizeof(T);
  constexpr int UNROLL_FACTOR = 128/sizeof(TN);
  constexpr int PEER_UNROLL = 2;

  T* sendPtr = (T*)ncclGetLocalPointer(sendwin, sendoffset);
  T* recvPtr = (T*)ncclGetLocalPointer(recvwin, recvoffset);

  // alignment check: can we use vectorized operations?
  bool canVectorize = (sizeof(TN) > sizeof(T)) &&
                      (reinterpret_cast<uintptr_t>(sendPtr) % sizeof(TN) == 0) &&
                      (reinterpret_cast<uintptr_t>(recvPtr) % sizeof(TN) == 0) &&
                      ((count * sizeof(T)) % sizeof(TN) == 0);

  if (canVectorize) {
    size_t vector_count = count / VECTOR_FACTOR;
    int elements_per_iteration = nthreads * UNROLL_FACTOR;

    // process aligned vectorized elements without bounds checks
    size_t aligned_vector_count = (vector_count / elements_per_iteration) * elements_per_iteration;
    for (size_t base_offset = tid; base_offset < aligned_vector_count; base_offset += elements_per_iteration) {
      // unroll a limited number of peers at a time
      for (int peerBase = 0; peerBase < lsaSize; peerBase += PEER_UNROLL) {
        int peersInGroup = min(PEER_UNROLL, lsaSize - peerBase);

        #pragma unroll
        for (int p = 0; p < peersInGroup; p++) {
          int lp = peerBase + p;
          TN* sendVecPtr = (TN*)(sendPtr + (size_t)(startLsa + lp) * count);
          TN* recvVecPtr = (TN*)((T*)ncclGetLsaPointer(recvwin, recvoffset, lp) + (size_t)worldRank * count);
          TN values[UNROLL_FACTOR];

          // split load/store into separate loops for better overlap and ILP
          #pragma unroll
          for (int i = 0; i < UNROLL_FACTOR; i++) {
            size_t offset = base_offset + i * nthreads;
            values[i] = sendVecPtr[offset];
          }
          #pragma unroll
          for (int i = 0; i < UNROLL_FACTOR; i++) {
            size_t offset = base_offset + i * nthreads;
            recvVecPtr[offset] = values[i];
          }
        }
      }
    }

    // handle remaining vectorized elements that didn't fit in aligned chunks
    for (size_t base_offset = aligned_vector_count + tid; base_offset < vector_count; base_offset += nthreads) {
      for (int lp = 0; lp < lsaSize; lp++) {
        TN* sendVecPtr = (TN*)(sendPtr + (size_t)(startLsa + lp) * count);
        TN* recvVecPtr = (TN*)((T*)ncclGetLsaPointer(recvwin, recvoffset, lp) + (size_t)worldRank * count);
        recvVecPtr[base_offset] = sendVecPtr[base_offset];
      }
    }
  } else {
    // simple scalar fallback for unaligned data (identical to simple kernel)
    for (size_t offset = tid; offset < count; offset += nthreads) {
      for (int lp = 0; lp < lsaSize; lp++) {
        T* recvPtr = (T*)ncclGetLsaPointer(recvwin, recvoffset, lp);
        recvPtr[(size_t)worldRank * count + offset] = sendPtr[(size_t)(startLsa + lp) * count + offset];
      }
    }
  }
}

// Device implementation #2 - Optimized NVL kernel
template <typename T>
__global__ void NvlAlltoAllKernelOptimized(ncclWindow_t sendwin, size_t sendoffset, ncclWindow_t recvwin, size_t recvoffset, size_t count, int root, struct ncclDevComm devComm) {
  ncclLsaBarrierSession<ncclCoopCta> bar { ncclCoopCta(), devComm, ncclTeamLsa(devComm), devComm.lsaBarrier, blockIdx.x };
  bar.sync(ncclCoopCta(), cuda::memory_order_acquire);

  int tid = threadIdx.x + blockDim.x * blockIdx.x;
  int nthreads = blockDim.x * gridDim.x;

  AlltoAllLsaVecImpl<T>(sendwin, sendoffset, recvwin, recvoffset, count, devComm.rank, 0, devComm.nRanks, tid, nthreads);

  bar.sync(ncclCoopCta(), cuda::memory_order_release);
}

#if NCCL_VERSION_CODE >= NCCL_VERSION(2,28,7) && defined(NCCL_OS_LINUX)
// None is canonical from 2.30.7; Relaxed is the only enumerator on earlier headers.
#if NCCL_VERSION_CODE >= NCCL_VERSION(2, 30, 7)
#define NCCL_TEST_GIN_FENCE_LEVEL ncclGinFenceLevel::None
#else
#define NCCL_TEST_GIN_FENCE_LEVEL ncclGinFenceLevel::Relaxed
#endif
// Device implementation #3 - GIN kernel requesting one context per GIN connection,
// but the kernel itself only uses context 0
template <typename T>
__global__ void GinAlltoAllKernel(ncclWindow_t sendwin, size_t sendoffset, ncclWindow_t recvwin, size_t recvoffset, size_t count, int root, struct ncclDevComm devComm) {
  int ginContext = 0;
  unsigned int signalIndex = blockIdx.x;
  ncclGin gin { devComm, ginContext };
  uint64_t signalValue = gin.readSignal(signalIndex);

#if NCCL_VERSION_CODE >= NCCL_VERSION(2, 30, 0)
  ncclGinBarrierSession<ncclCoopCta> bar { ncclCoopCta(), gin, ncclTeamTagWorld(), blockIdx.x };
#else
  ncclBarrierSession<ncclCoopCta> bar { ncclCoopCta(), ncclTeamTagWorld(), gin, blockIdx.x };
#endif
  bar.sync(ncclCoopCta(), cuda::memory_order_acquire, NCCL_TEST_GIN_FENCE_LEVEL);

  int tid = threadIdx.x + blockIdx.x * blockDim.x;
  int nthreads = blockDim.x * gridDim.x;

  /* send to all peers via GIN */
  const size_t size = count * sizeof(T);
  for (int r=tid; r<devComm.nRanks; r+=nthreads) {
    gin.put(ncclTeamWorld(devComm), r,
        recvwin, recvoffset + devComm.rank * size,
        sendwin, sendoffset + r * size,
        size,
#if NCCL_VERSION_CODE >= NCCL_VERSION(2, 30, 7)
        ncclGin_WeakSignalInc{signalIndex});
#else
        ncclGin_SignalInc{signalIndex});
#endif
  }

  int receivingCta = (devComm.rank % nthreads) / blockDim.x;
  if (blockIdx.x == receivingCta)
    gin.waitSignal(ncclCoopCta(), signalIndex, signalValue + devComm.nRanks);
  gin.flush(ncclCoopCta());
#if NCCL_VERSION_CODE < NCCL_VERSION(2, 30, 0)
  bar.sync(ncclCoopCta(), cuda::memory_order_release, NCCL_TEST_GIN_FENCE_LEVEL);
#endif
}

// Device implementation #4 - Hybrid kernel using LSA for local peers and GIN for remote ones
template <typename T>
__global__ void __launch_bounds__(512) HybridAlltoAllKernel(ncclWindow_t sendwin, size_t sendoffset, ncclWindow_t recvwin, size_t recvoffset, size_t count, int root, struct ncclDevComm devComm) {
  // Requires gridDim.x >= ginContextCount to use all available GIN devices
  int numCtx = min((int)gridDim.x, (int)devComm.ginContextCount);
  int myCtx = blockIdx.x % numCtx;
  int ctaInCtx = blockIdx.x / numCtx;
  int ctasPerCtx = ((int)gridDim.x - myCtx + numCtx - 1) / numCtx;

  ncclGin gin { devComm, myCtx };
  unsigned int signalIndex = myCtx;
  uint64_t signalValue = gin.readSignal(signalIndex);

  ncclBarrierSession<ncclCoopCta> bar { ncclCoopCta(), ncclTeamTagWorld(), gin, blockIdx.x };
  bar.sync(ncclCoopCta(), cuda::memory_order_acquire, NCCL_TEST_GIN_FENCE_LEVEL);

  int tid = threadIdx.x + blockIdx.x*blockDim.x;
  int nthreads = blockDim.x * gridDim.x;

  ncclTeam world = ncclTeamWorld(devComm);
  ncclTeam lsa = ncclTeamLsa(devComm);
  const int startLsa = world.rank - lsa.rank;
  const int lsaSize  = lsa.nRanks;

  /* handle remote peers (i.e., non-LSA) using GIN */
  const size_t size = count * sizeof(T);
  const size_t base = size / numCtx;
  const size_t rem = size % numCtx;
  const size_t sliceSize = base + (myCtx == numCtx - 1 ? rem : 0);
  const size_t sliceOff = (size_t)myCtx * base;

  int stride = ctasPerCtx * blockDim.x;
  if (sliceSize > 0) {
    for (int r = ctaInCtx + threadIdx.x*ctasPerCtx; r < world.nRanks; r += stride) {
      if (r < startLsa || r >= startLsa + lsaSize) {
        gin.put(world, r,
            recvwin, recvoffset + (size_t)world.rank * size + sliceOff,
            sendwin, sendoffset + (size_t)r * size + sliceOff,
            sliceSize,
#if NCCL_VERSION_CODE >= NCCL_VERSION(2, 30, 7)
            ncclGin_WeakSignalInc{signalIndex});
#else
            ncclGin_SignalInc{signalIndex});
#endif
      }
    }
  }

  /* handle local peers with LSA */
  AlltoAllLsaVecImpl<T>(sendwin, sendoffset, recvwin, recvoffset, count, world.rank, startLsa, lsaSize, tid, nthreads);

  // This context receives numRemotePeers increments; all CTAs sharing this context share
  // the same signal, so a single CTA waiter is enough
  int numRemotePeers = world.nRanks - lsa.nRanks;
  if (ctaInCtx == 0 && sliceSize > 0)
    gin.waitSignal(ncclCoopCta(), signalIndex, signalValue + numRemotePeers);
  gin.flush(ncclCoopCta());

  bar.sync(ncclCoopCta(), cuda::memory_order_release, NCCL_TEST_GIN_FENCE_LEVEL);
}

// Device implementation #5 - GIN kernel with one context per CTA
template <typename T>
__global__ void GinAlltoAllKernelMultiContext(ncclWindow_t sendwin, size_t sendoffset,
    ncclWindow_t recvwin, size_t recvoffset, size_t count, int root, struct ncclDevComm devComm) {
  // Requires gridDim.x >= ginContextCount to use all available GIN devices.
  int numCtx = min((int)gridDim.x, (int)devComm.ginContextCount);
  int myCtx = blockIdx.x % numCtx;
  int ctaInCtx = blockIdx.x / numCtx;
  int ctasPerCtx = ((int)gridDim.x - myCtx + numCtx - 1) / numCtx;

  ncclGin gin { devComm, myCtx };
  unsigned int signalIndex = myCtx;
  uint64_t signalValue = gin.readSignal(signalIndex);

#if NCCL_VERSION_CODE >= NCCL_VERSION(2, 30, 0)
  ncclGinBarrierSession<ncclCoopCta> bar { ncclCoopCta(), gin, ncclTeamTagWorld(), blockIdx.x };
#else
  ncclBarrierSession<ncclCoopCta> bar { ncclCoopCta(), ncclTeamTagWorld(), gin, blockIdx.x };
#endif
  bar.sync(ncclCoopCta(), cuda::memory_order_acquire, NCCL_TEST_GIN_FENCE_LEVEL);

  // Split each peer's block across the GIN contexts
  const size_t size = count * sizeof(T);
  const size_t base = size / numCtx;
  const size_t rem = size % numCtx;
  const size_t sliceSize = base + (myCtx == numCtx - 1 ? rem : 0);
  const size_t sliceOff = (size_t)myCtx * base;

  // nRanks work items distributed across the threads of this CTA
  int stride = ctasPerCtx * blockDim.x;
  if (sliceSize > 0) {
    for (int w = ctaInCtx + threadIdx.x * ctasPerCtx; w < devComm.nRanks; w += stride) {
      gin.put(ncclTeamWorld(devComm), w,
          recvwin, recvoffset + (size_t)devComm.rank*size + sliceOff,
          sendwin, sendoffset + (size_t)w*size + sliceOff,
          sliceSize,
#if NCCL_VERSION_CODE >= NCCL_VERSION(2, 30, 7)
          ncclGin_WeakSignalInc{signalIndex});
#else
          ncclGin_SignalInc{signalIndex});
#endif
    }
  }

  // This context receives nRanks increments; this CTA waits on them
  if (ctaInCtx == 0 && sliceSize > 0) {
    gin.waitSignal(ncclCoopCta(), signalIndex, signalValue + devComm.nRanks);
  }
  gin.flush(ncclCoopCta());
#if NCCL_VERSION_CODE < NCCL_VERSION(2, 30, 0)
  bar.sync(ncclCoopCta(), cuda::memory_order_release, NCCL_TEST_GIN_FENCE_LEVEL);
#endif
}
#endif
#endif

#if NCCL_VERSION_CODE >= NCCL_VERSION(2,29,0)
testResult_t AlltoAllRmaPut(void* sendWindow, size_t sendoffset, void* recvWindow, size_t recvoffset,
                            size_t count, ncclDataType_t type, ncclComm_t comm, cudaStream_t stream) {
  int rank, nranks;
  NCCLCHECK(ncclCommUserRank(comm, &rank));
  NCCLCHECK(ncclCommCount(comm, &nranks));

  ncclWindow_t sendWin = (ncclWindow_t)sendWindow;
  ncclWindow_t recvWin = (ncclWindow_t)recvWindow;

  void* sendPtr = NULL;
  void* recvPtr = NULL;
  NCCLCHECK(ncclWinGetUserPtr(comm, sendWin, &sendPtr));
  NCCLCHECK(ncclWinGetUserPtr(comm, recvWin, &recvPtr));

  size_t eltSize = wordSize(type);
  size_t chunkBytes = count * eltSize;
  const int nctx = rmaCtxCount;

  ncclWaitSignalDesc_t* waitDescs = (ncclWaitSignalDesc_t*)malloc(sizeof(ncclWaitSignalDesc_t) * nranks);
  if (waitDescs == NULL) {
    return testInternalError;
  }

  for (int i = 0; i < nranks; i++) {
    waitDescs[i].opCnt = 1;
    waitDescs[i].peer = i;
    waitDescs[i].sigIdx = i % NUM_RMA_SIG;
    waitDescs[i].ctx = (i + rank) % nctx;
  }

  NCCLCHECK(ncclGroupStart());
  for (int peer = 0; peer < nranks; peer++) {
    int targetRank = (rank + peer) % nranks;
    void* srcPtr = (char*)sendPtr + sendoffset + targetRank * chunkBytes;
    size_t dstOffset = recvoffset + rank * chunkBytes;

    NCCLCHECK(ncclPutSignal(srcPtr, count, type, targetRank,
                      recvWin, dstOffset, rank % NUM_RMA_SIG, (rank + targetRank) % nctx, 0, comm, stream));
  }
  NCCLCHECK(ncclGroupEnd());

  NCCLCHECK(ncclWaitSignal(nranks, waitDescs, comm, stream));
  free(waitDescs);
  return testSuccess;
}
#endif

testResult_t AlltoAllRunColl(void* sendbuff, size_t sendoffset, void* recvbuff, size_t recvoffset, size_t count, ncclDataType_t type, ncclRedOp_t op, int root, ncclComm_t comm, cudaStream_t stream, int deviceImpl) {
  if (deviceImpl == 0) {
    char* sptr = (char*)sendbuff + sendoffset;
    char* rptr = (char*)recvbuff + recvoffset;
#if NCCL_VERSION_CODE >= NCCL_VERSION(2,28,0)
    if (test_ncclVersion >= NCCL_VERSION(2,28,0)) {
      NCCLCHECK(ncclAlltoAll(sptr, rptr, count, type, comm, stream));
      return testSuccess;
    }
    // fall-through to send/recv implementation if ncclAlltoAll is not available
#endif
#if NCCL_VERSION_CODE >= NCCL_VERSION(2,7,0)
    int nRanks;
    NCCLCHECK(ncclCommCount(comm, &nRanks));
    size_t rankOffset = count * wordSize(type);
    NCCLCHECK(ncclGroupStart());
    for (int r=0; r<nRanks; r++) {
      NCCLCHECK(ncclSend(sptr+r*rankOffset, count, type, r, comm, stream));
      NCCLCHECK(ncclRecv(rptr+r*rankOffset, count, type, r, comm, stream));
    }
    NCCLCHECK(ncclGroupEnd());
#else
    printf("NCCL 2.7 or later is needed for alltoall. This test was compiled with %d.%d.\n", NCCL_MAJOR, NCCL_MINOR);
    return testNcclError;
#endif
  } else {
    switch(deviceImpl) {
#if NCCL_VERSION_CODE >= NCCL_VERSION(2,29,0)
      case HOST_RMA_IMPL:
        TESTCHECK(AlltoAllRmaPut(sendbuff, sendoffset, recvbuff, recvoffset, count, type, comm, stream));
        return testSuccess;
#endif
#if NCCL_VERSION_CODE >= NCCL_VERSION(2,28,0)
      case 1:
        TESTCHECK(testLaunchDeviceKernel(SPECIALIZE_KERNEL(NvlAlltoAllKernel, type, op), sendbuff, sendoffset, recvbuff, recvoffset, count, type, op, root, comm, stream));
        return testSuccess;
      case 2:
        TESTCHECK(testLaunchDeviceKernel(SPECIALIZE_KERNEL(NvlAlltoAllKernelOptimized, type, op), sendbuff, sendoffset, recvbuff, recvoffset, count, type, op, root, comm, stream));
        return testSuccess;
#endif
#if NCCL_VERSION_CODE >= NCCL_VERSION(2,28,7) && defined(NCCL_OS_LINUX)
      case 3:
        TESTCHECK(testLaunchDeviceKernel(SPECIALIZE_KERNEL(GinAlltoAllKernel, type, op), sendbuff, sendoffset, recvbuff, recvoffset, count, type, op, root, comm, stream));
        return testSuccess;
      case 4:
        TESTCHECK(testLaunchDeviceKernel(SPECIALIZE_KERNEL(HybridAlltoAllKernel, type, op), sendbuff, sendoffset, recvbuff, recvoffset, count, type, op, root, comm, stream));
        return testSuccess;
      case 5:
        TESTCHECK(testLaunchDeviceKernel(SPECIALIZE_KERNEL(GinAlltoAllKernelMultiContext, type, op), sendbuff, sendoffset, recvbuff, recvoffset, count, type, op, root, comm, stream));
        return testSuccess;
#endif
      default:
        return testNotImplemented;
    }
  }
  return testSuccess;
}

struct testColl alltoAllTest = {
  "AlltoAll",
  AlltoAllGetCollByteCount,
  AlltoAllInitData,
  AlltoAllGetBw,
  AlltoAllRunColl
};

void AlltoAllGetBuffSize(size_t *sendcount, size_t *recvcount, size_t count, int nranks) {
  size_t paramcount, sendInplaceOffset, recvInplaceOffset;
  AlltoAllGetCollByteCount(sendcount, recvcount, &paramcount, &sendInplaceOffset, &recvInplaceOffset, count, /*eltSize=*/1, nranks);
}

testResult_t AlltoAllRunTest(struct threadArgs* args, int root, ncclDataType_t type, const char* typeName, ncclRedOp_t op, const char* opName) {
  args->collTest = &alltoAllTest;
  ncclDataType_t *run_types;
  const char **run_typenames;
  int type_count;

  if ((int)type != -1) {
    type_count = 1;
    run_types = &type;
    run_typenames = &typeName;
  } else {
    type_count = test_typenum;
    run_types = test_types;
    run_typenames = test_typenames;
  }

  for (int i=0; i<type_count; i++) {
      TESTCHECK(TimeTest(args, run_types[i], run_typenames[i], (ncclRedOp_t)0, "none", -1));
  }
  return testSuccess;
}

NCCL_WEAK struct testEngine ncclTestEngine = {
  /* .getBuffSize = */ AlltoAllGetBuffSize,
  /* .runTest = */ AlltoAllRunTest,
#if NCCL_VERSION_CODE >= NCCL_VERSION(2,14,0)
  /* .initCommConfig = */ nullptr,
#endif
#if NCCL_VERSION_CODE >= NCCL_VERSION(2,28,0)
  /* .getDevCommRequirements = */ AlltoAllGetDevCommRequirements
#endif
};
