Add tilled version of puts and gets at wavefront level to the functional test suite
* Implemented tiled version of put*_wave and get*_wave functions * Maintain single kernel that supports both tiled and untiled versions * Disable IPC in the default RO build script
This commit is contained in:
@@ -27,17 +27,22 @@
|
||||
using namespace rocshmem;
|
||||
|
||||
/******************************************************************************
|
||||
* DEVICE TEST KERNEL
|
||||
* DEVICE TEST KERNELS
|
||||
*****************************************************************************/
|
||||
__global__ void WaveLevelPrimitiveTest(int loop, int skip, uint64_t *timer,
|
||||
char *s_buf, char *r_buf, int size,
|
||||
TestType type,
|
||||
ShmemContextType ctx_type) {
|
||||
TestType type, ShmemContextType ctx_type,
|
||||
int wf_size) {
|
||||
__shared__ roc_shmem_ctx_t ctx;
|
||||
roc_shmem_wg_init();
|
||||
roc_shmem_wg_ctx_create(ctx_type, &ctx);
|
||||
|
||||
uint64_t start;
|
||||
int wf_id = get_flat_block_id() / wf_size;
|
||||
int offset = size * get_flat_grid_id() * (get_flat_block_size() / wf_size);
|
||||
int idx = wf_id * size + offset;
|
||||
s_buf += idx;
|
||||
r_buf += idx;
|
||||
|
||||
for (int i = 0; i < loop + skip; i++) {
|
||||
if (i == skip) start = roc_shmem_timer();
|
||||
@@ -75,8 +80,10 @@ __global__ void WaveLevelPrimitiveTest(int loop, int skip, uint64_t *timer,
|
||||
*****************************************************************************/
|
||||
WaveLevelPrimitiveTester::WaveLevelPrimitiveTester(TesterArguments args)
|
||||
: Tester(args) {
|
||||
s_buf = (char *)roc_shmem_malloc(args.max_msg_size * args.wg_size);
|
||||
r_buf = (char *)roc_shmem_malloc(args.max_msg_size * args.wg_size);
|
||||
s_buf = (char *)roc_shmem_malloc(args.max_msg_size * args.num_wgs
|
||||
* num_warps);
|
||||
r_buf = (char *)roc_shmem_malloc(args.max_msg_size * args.num_wgs
|
||||
* num_warps);
|
||||
}
|
||||
|
||||
WaveLevelPrimitiveTester::~WaveLevelPrimitiveTester() {
|
||||
@@ -85,8 +92,8 @@ WaveLevelPrimitiveTester::~WaveLevelPrimitiveTester() {
|
||||
}
|
||||
|
||||
void WaveLevelPrimitiveTester::resetBuffers(uint64_t size) {
|
||||
memset(s_buf, '0', args.max_msg_size * args.wg_size);
|
||||
memset(r_buf, '1', args.max_msg_size * args.wg_size);
|
||||
memset(s_buf, '0', size * args.num_wgs * num_warps);
|
||||
memset(r_buf, '1', size * args.num_wgs * num_warps);
|
||||
}
|
||||
|
||||
void WaveLevelPrimitiveTester::launchKernel(dim3 gridSize, dim3 blockSize,
|
||||
@@ -95,10 +102,10 @@ void WaveLevelPrimitiveTester::launchKernel(dim3 gridSize, dim3 blockSize,
|
||||
|
||||
hipLaunchKernelGGL(WaveLevelPrimitiveTest, gridSize, blockSize, shared_bytes,
|
||||
stream, loop, args.skip, timer, s_buf, r_buf, size, _type,
|
||||
_shmem_context);
|
||||
_shmem_context, deviceProps.warpSize);
|
||||
|
||||
num_msgs = (loop + args.skip) * gridSize.x;
|
||||
num_timed_msgs = loop;
|
||||
num_timed_msgs = loop * gridSize.x;
|
||||
}
|
||||
|
||||
void WaveLevelPrimitiveTester::verifyResults(uint64_t size) {
|
||||
@@ -107,7 +114,7 @@ void WaveLevelPrimitiveTester::verifyResults(uint64_t size) {
|
||||
: 1;
|
||||
|
||||
if (args.myid == check_id) {
|
||||
for (int i = 0; i < size; i++) {
|
||||
for (int i = 0; i < size * args.num_wgs * num_warps; i++) {
|
||||
if (r_buf[i] != '0') {
|
||||
fprintf(stderr, "Data validation error at idx %d\n", i);
|
||||
fprintf(stderr, "Got %c, Expected %c \n", r_buf[i], '0');
|
||||
|
||||
Reference in New Issue
Block a user