diff --git a/hipamd/src/hip_stream.cpp b/hipamd/src/hip_stream.cpp index 22b509c439..4a9356c4c1 100644 --- a/hipamd/src/hip_stream.cpp +++ b/hipamd/src/hip_stream.cpp @@ -114,6 +114,25 @@ bool Stream::Create() { return result; } +// ================================================================================================ +bool isValid(hipStream_t& stream) { + // NULL stream is always valid + if (stream == nullptr) { + return true; + } + + if (hipStreamPerThread == stream) { + getStreamPerThread(stream); + } + + hip::Stream* s = reinterpret_cast(stream); + amd::ScopedLock lock(streamSetLock); + if (streamSet.find(s) == streamSet.end()) { + return false; + } + return true; +} + // ================================================================================================ amd::HostQueue* Stream::asHostQueue(bool skip_alloc) { if (queue_ != nullptr) { @@ -143,6 +162,12 @@ int Stream::DeviceId() const { } int Stream::DeviceId(const hipStream_t hStream) { + // Copying locally into non-const variable just to get const away + hipStream_t inputStream = hStream; + if (!hip::isValid(inputStream)) { + //return invalid device id + return -1; + } hip::Stream* s = reinterpret_cast(hStream); int deviceId = (s != nullptr)? s->DeviceId() : ihipGetDevice(); assert(deviceId >= 0 && deviceId < static_cast(g_devices.size())); @@ -175,25 +200,6 @@ void Stream::destroyAllStreams(int deviceId) { } } -// ================================================================================================ -bool isValid(hipStream_t& stream) { - // NULL stream is always valid - if (stream == nullptr) { - return true; - } - - if (hipStreamPerThread == stream) { - getStreamPerThread(stream); - } - - hip::Stream* s = reinterpret_cast(stream); - amd::ScopedLock lock(streamSetLock); - if (streamSet.find(s) == streamSet.end()) { - return false; - } - return true; -} - };// hip namespace // ================================================================================================