From 42a08627efbfd6e91c0a13e74230d57f7217197c Mon Sep 17 00:00:00 2001 From: sdashmiz Date: Mon, 3 Oct 2022 15:13:37 -0400 Subject: [PATCH] SWDEV-360031 - check for stream capture finish. - stream capture should be done before any sync APIs. Signed-off-by: sdashmiz Change-Id: I3d65f67ee68777be71f97f48d460ccaefdd4e1af [ROCm/clr commit: be966acb0cc9c61760369e0649b9f709a0481ddc] --- .../clr/hipamd/src/hip_device_runtime.cpp | 4 ++++ projects/clr/hipamd/src/hip_event.cpp | 3 +++ projects/clr/hipamd/src/hip_graph.cpp | 10 ++++++++ projects/clr/hipamd/src/hip_internal.hpp | 4 ++++ projects/clr/hipamd/src/hip_stream.cpp | 23 ++++++++++++++++++- 5 files changed, 43 insertions(+), 1 deletion(-) diff --git a/projects/clr/hipamd/src/hip_device_runtime.cpp b/projects/clr/hipamd/src/hip_device_runtime.cpp index 6e8dfdb4f0..def9bcb0ef 100644 --- a/projects/clr/hipamd/src/hip_device_runtime.cpp +++ b/projects/clr/hipamd/src/hip_device_runtime.cpp @@ -518,6 +518,10 @@ hipError_t hipDeviceSynchronize ( void ) { HIP_RETURN(hipErrorOutOfMemory); } + if (hip::Stream::StreamCaptureOngoing() == true) { + HIP_RETURN(hipErrorStreamCaptureUnsupported); + } + queue->finish(); hip::Stream::syncNonBlockingStreams(hip::getCurrentDevice()->deviceId()); diff --git a/projects/clr/hipamd/src/hip_event.cpp b/projects/clr/hipamd/src/hip_event.cpp index 7f384eff65..f556cabe9d 100644 --- a/projects/clr/hipamd/src/hip_event.cpp +++ b/projects/clr/hipamd/src/hip_event.cpp @@ -404,6 +404,9 @@ hipError_t hipEventSynchronize(hipEvent_t event) { HIP_RETURN(hipErrorInvalidHandle); } + if (hip::Stream::StreamCaptureOngoing() == true) { + HIP_RETURN(hipErrorStreamCaptureUnsupported); + } hip::Event* e = reinterpret_cast(event); HIP_RETURN(e->synchronize()); } diff --git a/projects/clr/hipamd/src/hip_graph.cpp b/projects/clr/hipamd/src/hip_graph.cpp index a41d3b5d3c..529b155840 100644 --- a/projects/clr/hipamd/src/hip_graph.cpp +++ b/projects/clr/hipamd/src/hip_graph.cpp @@ -29,6 +29,8 @@ std::vector g_captureStreams; amd::Monitor g_captureStreamsLock{"StreamCaptureGlobalList"}; +static amd::Monitor g_streamSetLock{"StreamCaptureset"}; +std::unordered_set g_allCapturingStreams; inline hipError_t ihipGraphAddNode(hipGraphNode_t graphNode, hipGraph_t graph, const hipGraphNode_t* pDependencies, size_t numDependencies) { @@ -959,6 +961,10 @@ hipError_t hipStreamBeginCapture_common(hipStream_t stream, hipStreamCaptureMode amd::ScopedLock lock(g_captureStreamsLock); g_captureStreams.push_back(s); } + { + amd::ScopedLock lock(g_streamSetLock); + g_allCapturingStreams.insert(s); + } return hipSuccess; } @@ -1010,6 +1016,10 @@ hipError_t hipStreamEndCapture_common(hipStream_t stream, hipGraph_t* pGraph) { *pGraph = nullptr; return hipErrorStreamCaptureInvalidated; } + { + amd::ScopedLock lock(g_streamSetLock); + g_allCapturingStreams.erase(std::find(g_allCapturingStreams.begin(), g_allCapturingStreams.end(), s)); + } // check if all parallel streams have joined // Nodes that are removed from the dependency set via API hipStreamUpdateCaptureDependencies do // not result in hipErrorStreamCaptureUnjoined diff --git a/projects/clr/hipamd/src/hip_internal.hpp b/projects/clr/hipamd/src/hip_internal.hpp index 1a88308651..be90fe5b1c 100644 --- a/projects/clr/hipamd/src/hip_internal.hpp +++ b/projects/clr/hipamd/src/hip_internal.hpp @@ -298,6 +298,9 @@ namespace hip { /// Destroy all streams on a given device static void destroyAllStreams(int deviceId); + /// Check Stream Capture status to make sure it is done + static bool StreamCaptureOngoing(void); + /// Returns capture status of the current stream hipStreamCaptureStatus GetCaptureStatus() const { return captureStatus_; } /// Returns capture mode of the current stream @@ -566,4 +569,5 @@ constexpr bool kMarkerDisableFlush = true; //!< Avoids command batch flush in extern std::vector g_captureStreams; extern amd::Monitor g_captureStreamsLock; +extern std::unordered_set g_allCapturingStreams; #endif // HIP_SRC_HIP_INTERNAL_H diff --git a/projects/clr/hipamd/src/hip_stream.cpp b/projects/clr/hipamd/src/hip_stream.cpp index 36fad7ddfc..3d1e9168dc 100644 --- a/projects/clr/hipamd/src/hip_stream.cpp +++ b/projects/clr/hipamd/src/hip_stream.cpp @@ -211,6 +211,10 @@ void Stream::destroyAllStreams(int deviceId) { } } +bool Stream::StreamCaptureOngoing(void) { + return (g_allCapturingStreams.empty() == true) ? false : true; +} + };// hip namespace // ================================================================================================ @@ -442,6 +446,12 @@ hipError_t hipStreamSynchronize_common(hipStream_t stream) { if (!hip::isValid(stream)) { HIP_RETURN(hipErrorContextIsDestroyed); } + if (stream != nullptr) { + // If still capturing return error + if (hip::Stream::StreamCaptureOngoing() == true) { + HIP_RETURN(hipErrorStreamCaptureUnsupported); + } + } // Wait for the current host queue hip::getQueue(stream)->finish(); return hipSuccess; @@ -524,6 +534,12 @@ hipError_t hipStreamWaitEvent_common(hipStream_t stream, hipEvent_t event, unsig return hipErrorContextIsDestroyed; } + if (stream != nullptr) { + // If still capturing return error + if (hip::Stream::StreamCaptureOngoing() == true) { + HIP_RETURN(hipErrorStreamCaptureIsolation); + } + } hip::Event* e = reinterpret_cast(event); return e->streamWait(stream, flags); } @@ -546,7 +562,12 @@ hipError_t hipStreamQuery_common(hipStream_t stream) { if (!hip::isValid(stream)) { return hipErrorContextIsDestroyed; } - + if (stream != nullptr) { + // If still capturing return error + if (hip::Stream::StreamCaptureOngoing() == true) { + HIP_RETURN(hipErrorStreamCaptureUnsupported); + } + } amd::HostQueue* hostQueue = hip::getQueue(stream); amd::Command* command = hostQueue->getLastQueuedCommand(true);