Added alltoallv test and optional args variable on collective args (#514)

* Added alltoallv test and optional args variable on collective args
Этот коммит содержится в:
akolliasAMD
2022-03-18 13:55:11 -04:00
коммит произвёл GitHub
родитель a04da71647
Коммит 65ea3d80db
13 изменённых файлов: 284 добавлений и 154 удалений
+21 -26
Просмотреть файл
@@ -14,20 +14,17 @@ namespace RcclUnitTesting
int const deviceId,
ncclFunc_t const funcType,
ncclDataType_t const dataType,
ncclRedOp_t const redOp,
int const root,
size_t const numInputElements,
size_t const numOutputElements,
ScalarTransport const scalarTransport,
int const scalarMode)
OptionalColArgs const &optionalColArgs)
{
// Free scalar based on previous scalarMode
if (scalarMode != -1)
if (optionalColArgs.scalarMode != -1)
{
if (this->localScalar.ptr != nullptr)
{
if (this->scalarMode == 0) this->localScalar.FreeGpuMem();
if (this->scalarMode == 1) hipHostFree(this->localScalar.ptr);
if (this->options.scalarMode == 0) this->localScalar.FreeGpuMem();
if (this->options.scalarMode == 1) hipHostFree(this->localScalar.ptr);
}
}
@@ -36,26 +33,23 @@ namespace RcclUnitTesting
this->deviceId = deviceId;
this->funcType = funcType;
this->dataType = dataType;
this->redOp = redOp;
this->root = root;
this->numInputElements = numInputElements;
this->numOutputElements = numOutputElements;
this->scalarTransport = scalarTransport;
this->scalarMode = scalarMode;
this->options = optionalColArgs;
if (scalarMode != -1)
if (this->options.scalarMode != -1)
{
size_t const numBytes = DataTypeToBytes(dataType);
if (scalarMode == ncclScalarDevice)
if (this->options.scalarMode == ncclScalarDevice)
{
CHECK_CALL(this->localScalar.AllocateGpuMem(numBytes));
CHECK_HIP(hipMemcpy(this->localScalar.ptr, scalarTransport.ptr + (globalRank * numBytes),
CHECK_HIP(hipMemcpy(this->localScalar.ptr, optionalColArgs.scalarTransport.ptr + (globalRank * numBytes),
numBytes, hipMemcpyHostToDevice));
}
else if (scalarMode == ncclScalarHostImmediate)
else if (this->options.scalarMode == ncclScalarHostImmediate)
{
CHECK_HIP(hipHostMalloc(&this->localScalar.ptr, numBytes, 0));
memcpy(this->localScalar.ptr, scalarTransport.ptr + (globalRank * numBytes), numBytes);
memcpy(this->localScalar.ptr, optionalColArgs.scalarTransport.ptr + (globalRank * numBytes), numBytes);
}
}
return TEST_SUCCESS;
@@ -116,8 +110,8 @@ namespace RcclUnitTesting
ErrCode CollectiveArgs::ValidateResults()
{
// Ignore non-root outputs for collectives with a root
if (CollectiveArgs::UsesRoot(this->funcType) && this->root != this->globalRank) return TEST_SUCCESS;
if (CollectiveArgs::UsesRoot(this->funcType) && this->options.root != this->globalRank) return TEST_SUCCESS;
if (this->funcType == ncclCollSend) return TEST_SUCCESS; // on the send receive pair only recv needs to be checked
size_t const numOutputBytes = (this->numOutputElements * DataTypeToBytes(this->dataType));
CHECK_HIP(hipMemcpy(this->outputCpu.ptr, this->outputGpu.ptr, numOutputBytes, hipMemcpyDeviceToHost));
@@ -153,8 +147,9 @@ namespace RcclUnitTesting
if (this->localScalar.ptr != nullptr)
{
if (this->scalarMode == 0) this->localScalar.FreeGpuMem();
if (this->scalarMode == 1) CHECK_HIP(hipHostFree(this->localScalar.ptr));
if (this->options.scalarMode == 0) this->localScalar.FreeGpuMem();
if (this->options.scalarMode == 1) CHECK_HIP(hipHostFree(this->localScalar.ptr));
this->localScalar.Attach(nullptr);
}
return TEST_SUCCESS;
}
@@ -174,6 +169,7 @@ namespace RcclUnitTesting
case ncclCollGather: ss << "ncclGather"; break;
case ncclCollScatter: ss << "ncclScatter"; break;
case ncclCollAllToAll: ss << "ncclAllToAll"; break;
case ncclCollAllToAllv: ss << "ncclAllToAllv"; break;
case ncclCollSend: ss << "ncclSend"; break;
case ncclCollRecv: ss << "ncclRecv"; break;
default: ss << "[Unknown]"; break;
@@ -184,9 +180,9 @@ namespace RcclUnitTesting
this->funcType == ncclCollReduceScatter ||
this->funcType == ncclCollAllReduce)
{
if (this->redOp < ncclNumOps)
if (this->options.redOp < ncclNumOps)
{
ss << ncclRedOpNames[this->redOp] << " ";
ss << ncclRedOpNames[this->options.redOp] << " ";
}
else
{
@@ -215,13 +211,13 @@ namespace RcclUnitTesting
this->funcType == ncclCollGather ||
this->funcType == ncclCollScatter)
{
ss << "Root " << this->root << " ";
ss << "Root " << this->options.root << " ";
}
if (this->funcType == ncclCollSend ||
this->funcType == ncclCollRecv)
{
ss << "Peer " << this->root << " ";
ss << "Peer " << this->options.root << " ";
}
ss << "#In: " << this->numInputElements;
@@ -277,7 +273,6 @@ namespace RcclUnitTesting
return (funcType == ncclCollBroadcast ||
funcType == ncclCollReduce ||
funcType == ncclCollGather ||
funcType == ncclCollScatter ||
funcType == ncclCollSend); // this is incorrect but it works because in Send root is not root it is the peer
funcType == ncclCollScatter);
}
}