Add host API for *_on_stream operations (#340)
* Add functional test for barrier_all_on_stream
* Add rocshmem_barrier_all_on_stream support for GDA and RO backends
Implements rocshmem_barrier_all_on_stream operation for
GPU Direct Access and Reverse Offload backends.
Previously, rocshmem_barrier_all_on_stream was only supported for IPC backend.
* Add functional test for rocshmem_broadcastmem_on_stream
* Add host-side rocshmem_broadcastmem_on_stream API
Implement stream-based broadcast collective operation
- Add rocshmem_broadcastmem_on_stream host API and kernel implementation
- Add functional test TeamBroadcastmemOnStreamTester with multi-stream
support and correctness verification
- Use per-workgroup contexts to avoid contention across parallel streams
API:
rocshmem_broadcastmem_on_stream(team, dest, source, nelems, pe_root, stream)
* Add functional test for rocshmem_getmem_on_stream
* Add host-side rocshmem_getmem_on_stream API
Implement stream-based point-to-point RMA get operation
- Add rocshmem_getmem_on_stream host API and kernel implementation
- Support for asynchronous getmem operations on HIP streams
- Add backend support for GDA, RO, and IPC contexts
- Use work-group collective getmem for efficient memory transfer
API:
rocshmem_getmem_on_stream(dest, source, nelems, pe, stream)
(AI Assist)
* Add host-side rocshmem_putmem_on_stream API
- Add rocshmem_putmem_on_stream for asynchronous remote writes
- Support for concurrent RMA operations on HIP streams
- Add backend support for GDA, RO, and IPC contexts
- Use work-group device collective operation
API:
rocshmem_putmem_on_stream(dest, source, bytes, pe, stream)
(AI Assist)
* Add functional test for rocshmem_putmem_on_stream
* Add host-side rocshmem_putmem_signal_on_stream API
Enables asynchronous putmem operations with signaling on HIP streams.
The implementation includes:
- Kernel wrapper rocshmem_putmem_signal_kernel
- Host interface putmem_signal_on_stream method
- Context layer support across all backends (IPC, GDA, RO)
- Public API
Function signature:
void rocshmem_putmem_signal_on_stream(void *dest, const void *source,
size_t bytes, uint64_t *sig_addr,
uint64_t signal, int sig_op,
int pe, hipStream_t stream);
* Add functional test for rocshmem_putmem_signal_on_stream
* Add host-side rocshmem_signal_wait_until_on_stream API
Enables asynchronous signal wait operations on HIP streams.
The implementation includes:
- Kernel wrapper rocshmem_signal_wait_until_kernel
- Host interface signal_wait_until_on_stream method
- Context layer support across all backends (IPC, GDA, RO)
- Native uint64_t support in wait_until API (generated from P2P_SYNC.py)
Function signature:
void rocshmem_signal_wait_until_on_stream(uint64_t *sig_addr, int cmp,
uint64_t cmp_value,
hipStream_t stream);
(AI Assist)
* Add functional test for rocshmem_signal_wait_until_on_stream
* Add documentation for stream API functions
This commit adds API documentation for the following host-side
stream functions:
- rocshmem_barrier_all_on_stream (collective routines)
- rocshmem_broadcastmem_on_stream (collective routines)
- rocshmem_getmem_on_stream (RMA operations)
- rocshmem_putmem_on_stream (RMA operations)
- rocshmem_putmem_signal_on_stream (signaling operations)
- rocshmem_signal_wait_until_on_stream (point-to-point sync)
The documentation includes function signatures, parameter descriptions,
and detailed explanations of asynchronous behavior and stream handling.
(AI Assist)
* Rename "bytes" -> "nelems"
* Add "_TEST_" to the variables used in tests
* Remove incorrect hipStreamDefault usage
hipStreamDefault is not a default stream. This is a flag.
If stream == nullptr, then just pass it to kernel. It will launch the kernel on the default stream
[ROCm/rocshmem commit: d0c8380650]
Este commit está contenido en:
@@ -36,7 +36,12 @@
|
||||
#include "amo_standard_tester.hpp"
|
||||
#include "default_ctx_primitive_tester.hpp"
|
||||
#include "barrier_all_tester.hpp"
|
||||
#include "barrier_all_on_stream_tester.hpp"
|
||||
#include "empty_tester.hpp"
|
||||
#include "getmem_on_stream_tester.hpp"
|
||||
#include "putmem_on_stream_tester.hpp"
|
||||
#include "putmem_signal_on_stream_tester.hpp"
|
||||
#include "signal_wait_until_on_stream_tester.hpp"
|
||||
#include "ping_all_tester.hpp"
|
||||
#include "ping_pong_tester.hpp"
|
||||
#include "primitive_mr_tester.hpp"
|
||||
@@ -48,6 +53,7 @@
|
||||
#include "team_sync_tester.hpp"
|
||||
#include "team_alltoall_tester.hpp"
|
||||
#include "team_alltoallmem_on_stream_tester.hpp"
|
||||
#include "team_broadcastmem_on_stream_tester.hpp"
|
||||
#include "team_barrier_tester.hpp"
|
||||
#include "team_broadcast_tester.hpp"
|
||||
#include "team_ctx_infra_tester.hpp"
|
||||
@@ -233,6 +239,36 @@ std::vector<Tester*> Tester::create(TesterArguments args) {
|
||||
std::cout << "Alltoallmem_On_Stream ###" << std::endl;
|
||||
testers.push_back(new TeamAlltoallmemOnStreamTester(args));
|
||||
return testers;
|
||||
case BarrierAllOnStreamTestType:
|
||||
if (rank == 0)
|
||||
std::cout << "Barrier_All_On_Stream ###" << std::endl;
|
||||
testers.push_back(new BarrierAllOnStreamTester(args));
|
||||
return testers;
|
||||
case TeamBroadcastmemOnStreamTestType:
|
||||
if (rank == 0)
|
||||
std::cout << "Broadcastmem_On_Stream ###" << std::endl;
|
||||
testers.push_back(new TeamBroadcastmemOnStreamTester(args));
|
||||
return testers;
|
||||
case GetmemOnStreamTestType:
|
||||
if (rank == 0)
|
||||
std::cout << "Getmem_On_Stream ###" << std::endl;
|
||||
testers.push_back(new GetmemOnStreamTester(args));
|
||||
return testers;
|
||||
case PutmemOnStreamTestType:
|
||||
if (rank == 0)
|
||||
std::cout << "Putmem_On_Stream ###" << std::endl;
|
||||
testers.push_back(new PutmemOnStreamTester(args));
|
||||
return testers;
|
||||
case PutmemSignalOnStreamTestType:
|
||||
if (rank == 0)
|
||||
std::cout << "Putmem_Signal_On_Stream ###" << std::endl;
|
||||
testers.push_back(new PutmemSignalOnStreamTester(args));
|
||||
return testers;
|
||||
case SignalWaitUntilOnStreamTestType:
|
||||
if (rank == 0)
|
||||
std::cout << "Signal_Wait_Until_On_Stream ###" << std::endl;
|
||||
testers.push_back(new SignalWaitUntilOnStreamTester(args));
|
||||
return testers;
|
||||
case TeamFCollectTestType:
|
||||
if (rank == 0) {
|
||||
std::cout << "Fcollect Test ###" << std::endl;
|
||||
@@ -569,30 +605,50 @@ void Tester::execute() {
|
||||
}
|
||||
|
||||
bool Tester::peLaunchesKernel() {
|
||||
bool is_launcher;
|
||||
|
||||
/**
|
||||
* The PE assigned 0 is always active in these tests.
|
||||
*/
|
||||
is_launcher = args.myid == 0;
|
||||
bool is_launcher = (args.myid == 0);
|
||||
|
||||
/**
|
||||
* Some test types are active on both sides.
|
||||
*/
|
||||
is_launcher = is_launcher || (_type == TeamReductionTestType) ||
|
||||
(_type == TeamBroadcastTestType) || (_type == TeamCtxInfraTestType) ||
|
||||
(_type == TeamCtxInfraTestSingleType) || (_type == TeamCtxInfraTestBlockType) ||
|
||||
(_type == TeamCtxInfraTestOddEvenType) ||
|
||||
(_type == TeamAllToAllTestType) || (_type == TeamFCollectTestType) ||
|
||||
(_type == PingPongTestType) || (_type == BarrierAllTestType) ||
|
||||
(_type == WAVEBarrierAllTestType) || (_type == WGBarrierAllTestType) ||
|
||||
(_type == TeamSyncTestType) || (_type == TeamWAVESyncTestType) ||
|
||||
(_type == TeamWGSyncTestType) || (_type == SyncAllTestType) ||
|
||||
(_type == WAVESyncAllTestType) || (_type == WGSyncAllTestType) ||
|
||||
(_type == RandomAccessTestType) || (_type == PingAllTestType) ||
|
||||
(_type == TeamBarrierTestType) || (_type == TeamWAVEBarrierTestType) ||
|
||||
(_type == TeamWGBarrierTestType) ||
|
||||
(_type == TeamAlltoallmemOnStreamTestType);
|
||||
switch (_type) {
|
||||
case TeamReductionTestType:
|
||||
case TeamBroadcastTestType:
|
||||
case TeamCtxInfraTestType:
|
||||
case TeamCtxInfraTestSingleType:
|
||||
case TeamCtxInfraTestBlockType:
|
||||
case TeamCtxInfraTestOddEvenType:
|
||||
case TeamAllToAllTestType:
|
||||
case TeamFCollectTestType:
|
||||
case PingPongTestType:
|
||||
case BarrierAllTestType:
|
||||
case WAVEBarrierAllTestType:
|
||||
case WGBarrierAllTestType:
|
||||
case TeamSyncTestType:
|
||||
case TeamWAVESyncTestType:
|
||||
case TeamWGSyncTestType:
|
||||
case SyncAllTestType:
|
||||
case WAVESyncAllTestType:
|
||||
case WGSyncAllTestType:
|
||||
case RandomAccessTestType:
|
||||
case PingAllTestType:
|
||||
case TeamBarrierTestType:
|
||||
case TeamWAVEBarrierTestType:
|
||||
case TeamWGBarrierTestType:
|
||||
case TeamAlltoallmemOnStreamTestType:
|
||||
case BarrierAllOnStreamTestType:
|
||||
case TeamBroadcastmemOnStreamTestType:
|
||||
case GetmemOnStreamTestType:
|
||||
case PutmemOnStreamTestType:
|
||||
case PutmemSignalOnStreamTestType:
|
||||
case SignalWaitUntilOnStreamTestType:
|
||||
is_launcher = true;
|
||||
break;
|
||||
default:
|
||||
break;
|
||||
}
|
||||
|
||||
return is_launcher;
|
||||
}
|
||||
|
||||
Referencia en una nueva incidencia
Block a user