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:
Anatolii Rozanov
2025-12-09 15:55:46 +01:00
cometido por GitHub
padre b9c172de16
commit f98c72d627
Se han modificado 39 ficheros con 2649 adiciones y 49 borrados
@@ -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;
}