[llvm] [Offload][Lang] Add internal StreamTy (PR #216374)

Sophia Herrmann via llvm-commits llvm-commits at lists.llvm.org
Thu Aug 27 13:56:16 PDT 2026


https://github.com/jellytabby updated https://github.com/llvm/llvm-project/pull/216374

>From 9c0f5dc07a9ffa1ec0c26b4e5f184dfe3fd77f17 Mon Sep 17 00:00:00 2001
From: Sophia Herrmann <herrmann15 at llnl.gov>
Date: Thu, 13 Aug 2026 16:03:12 -0700
Subject: [PATCH] add interal StreamTy

---
 .../languages/kernel/include/LanguageUtils.h  |  24 +-
 .../languages/kernel/include/OffloadErrors.h  |  38 +++
 offload/languages/kernel/include/State.h      | 115 +++++--
 offload/languages/kernel/include/Stream.h     |  33 ++
 .../languages/kernel/src/LanguageErrors.cpp   |   7 +-
 .../languages/kernel/src/LanguageLaunch.cpp   |  33 +-
 .../kernel/src/LanguageRegistration.cpp       |  19 +-
 .../languages/kernel/src/LanguageRuntime.cpp  |  58 ++--
 offload/languages/kernel/src/State.cpp        | 323 +++++++++++++-----
 offload/test/offloading/CUDA/stream_api.cu    |  40 ++-
 offload/test/offloading/HIP/stream_api.hip    |  38 ++-
 11 files changed, 533 insertions(+), 195 deletions(-)
 create mode 100644 offload/languages/kernel/include/OffloadErrors.h
 create mode 100644 offload/languages/kernel/include/Stream.h

diff --git a/offload/languages/kernel/include/LanguageUtils.h b/offload/languages/kernel/include/LanguageUtils.h
index 9726d0599a9f4..1116f9e550f0d 100644
--- a/offload/languages/kernel/include/LanguageUtils.h
+++ b/offload/languages/kernel/include/LanguageUtils.h
@@ -12,6 +12,10 @@
 #include "LanguageRuntime.h"
 #include "OffloadAPI.h"
 #include "State.h"
+#include "Stream.h"
+
+namespace llvm {
+namespace offload {
 
 /// Convert an ol_result_t to the active language's Error_t.
 static inline Error_t convertResult(ol_result_t Result) {
@@ -39,8 +43,7 @@ static inline Error_t convertResult(ol_result_t Result) {
 /// Set the last error for the current thread and return it.
 static inline Error_t setLastError(Error_t Error) {
   // TODO: find a more efficient way to set last error
-  return static_cast<Error_t>(
-      llvm::offload::ThreadStateTy::setLastError(Error));
+  return static_cast<Error_t>(ThreadStateTy::get().setLastError(Error));
 }
 
 /// Convert an ol_result_t to the active language's Error_t and set it as the
@@ -49,13 +52,28 @@ static inline Error_t convertAndSetLastError(ol_result_t Result) {
   return setLastError(convertResult(Result));
 }
 
+/// Convert between the language-facing opaque stream and the internal stream.
+static inline Stream_t toLanguageStream(StreamTy *Stream) {
+  return reinterpret_cast<Stream_t>(Stream);
+}
+
+static inline StreamTy *toInternalStream(Stream_t Stream) {
+  return reinterpret_cast<StreamTy *>(Stream);
+}
+
 /// Convert a Stream_t to an ol_queue_handle_t.
 static inline Error_t getQueueFromStream(Stream_t Stream,
                                          ol_queue_handle_t *Queue) {
   if (!Stream)
     return ErrorInvalidValue;
-  *Queue = reinterpret_cast<ol_queue_handle_t>(Stream);
+
+  //TODO: add proper DEBUG/assert guarded checks
+  StreamTy *InternalStream = toInternalStream(Stream);
+  *Queue = InternalStream->Queue;
   return Success;
 }
 
+} // namespace offload
+} // namespace llvm
+
 #endif // LLVM_OFFLOAD_LANGUAGES_KERNEL_INCLUDE_LANGUAGE_UTILS_H
diff --git a/offload/languages/kernel/include/OffloadErrors.h b/offload/languages/kernel/include/OffloadErrors.h
new file mode 100644
index 0000000000000..fe7477bab9ac4
--- /dev/null
+++ b/offload/languages/kernel/include/OffloadErrors.h
@@ -0,0 +1,38 @@
+//===-- OffloadErrors.h - Kernel language offload errors ------------------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+
+#ifndef LLVM_OFFLOAD_LANGUAGES_KERNEL_INCLUDE_OFFLOAD_ERRORS_H
+#define LLVM_OFFLOAD_LANGUAGES_KERNEL_INCLUDE_OFFLOAD_ERRORS_H
+
+#include "OffloadAPI.h"
+
+namespace llvm {
+namespace offload {
+
+inline constexpr ol_error_struct_t InvalidKernelError = {
+    OL_ERRC_INVALID_NULL_HANDLE, "kernel is not registered"};
+
+inline constexpr ol_error_struct_t InvalidDeviceError = {OL_ERRC_INVALID_DEVICE,
+                                                         "invalid device"};
+
+inline constexpr ol_error_struct_t InvalidArgumentError = {
+    OL_ERRC_INVALID_ARGUMENT, "invalid argument"};
+
+inline constexpr ol_error_struct_t InvalidConfigurationError = {
+    OL_ERRC_INVALID_SIZE, "invalid kernel launch configuration"};
+
+inline constexpr ol_error_struct_t InvalidNullPointerError = {
+    OL_ERRC_INVALID_NULL_POINTER, "invalid null pointer"};
+
+inline constexpr ol_error_struct_t InvalidStreamError = {OL_ERRC_INVALID_QUEUE,
+                                                         "invalid stream"};
+
+} // namespace offload
+} // namespace llvm
+
+#endif // LLVM_OFFLOAD_LANGUAGES_KERNEL_INCLUDE_OFFLOAD_ERRORS_H
diff --git a/offload/languages/kernel/include/State.h b/offload/languages/kernel/include/State.h
index 7315b2e710df5..a3c8712723a7a 100644
--- a/offload/languages/kernel/include/State.h
+++ b/offload/languages/kernel/include/State.h
@@ -10,13 +10,16 @@
 #define LLVM_OFFLOAD_LANGUAGES_KERNEL_INCLUDE_STATE_H
 
 #include "OffloadAPI.h"
+#include "Stream.h"
 #include "Types.h"
 
 #include "llvm/ADT/ArrayRef.h"
 #include "llvm/ADT/DenseMap.h"
+#include "llvm/ADT/SmallPtrSet.h"
 #include "llvm/ADT/SmallVector.h"
 #include "llvm/Support/raw_ostream.h"
 #include <cstdint>
+#include <mutex>
 
 #define CHECK_FATAL(ResultExpr, ...)                                           \
   do {                                                                         \
@@ -39,6 +42,9 @@
 namespace llvm {
 namespace offload {
 
+static constexpr unsigned AssumedDeviceCount = 8;
+static constexpr unsigned AssumedStreamCount = 8;
+
 /// Opaque host-side key used to identify a registered kernel.
 ///
 /// This is the address emitted in the offload entry table for the kernel
@@ -51,41 +57,43 @@ using KernelIDTy = const void *;
 struct ThreadStateTy {
   ~ThreadStateTy();
 
-  /// Return the default queue for the current host thread
-  static ol_queue_handle_t getDefaultQueue();
+  /// Return the thread-local state for the current host thread.
+  static ThreadStateTy &get();
+
+  /// Return the default queue for the current host thread.
+  ol_queue_handle_t getDefaultQueue();
+
+  /// Return the default stream for the current host thread and device.
+  StreamTy *getDefaultStream();
 
   /// Return the thread-local default device, or the first discovered device.
-  static ol_device_handle_t getDefaultDevice();
+  ol_device_handle_t getDefaultDevice();
 
   /// Return the thread-local default device and write its number to \p
   /// DeviceNo.
-  static ol_device_handle_t getDevice(int *DeviceNo);
+  ol_device_handle_t getDevice(int *DeviceNo);
 
   /// Set the thread-local default device by device number.
   ///
   /// \returns the selected device, or nullptr if \p DeviceNo is invalid.
-  static ol_device_handle_t setDefaultDevice(int DeviceNo);
+  ol_device_handle_t setDefaultDevice(int DeviceNo);
 
   /// Return the last language-runtime error code for this thread.
-  static uint32_t getLastError();
+  uint32_t getLastError();
 
   /// Set the last language-runtime error code for this thread.
-  static uint32_t setLastError(uint32_t Error);
+  uint32_t setLastError(uint32_t Error);
 
   /// Return the pending kernel launch configuration for this thread.
-  static CallConfigurationTy &getCallConfiguration();
-
-  /// Set the thread-local default device to \p Device and recreate its queue.
-  static void setDefaultDevice(ol_device_handle_t Device);
+  CallConfigurationTy &getCallConfiguration();
 
 private:
-  static ThreadStateTy &get();
-
-  void createDefaultQueue(ol_device_handle_t Device);
+  StreamTy *getOrCreateDefaultStream(ol_device_handle_t Device);
+  void destroyDefaultStreams();
 
   int DefaultDevice = 0;
   uint32_t LastError = 0;
-  ol_queue_handle_t DefaultQueue = nullptr;
+  DenseMap<ol_device_handle_t, StreamTy *> PerThreadDeviceDefaultStreamMap;
 
   CallConfigurationTy CC = {};
 
@@ -101,69 +109,102 @@ struct StateTy {
 
   friend struct ThreadStateTy;
 
+  /// Return the process-wide state singleton.
+  static StateTy &get();
+
+  /// Return the process-wide state singleton if it has been initialized.
+  static StateTy *tryGet();
+
   /// Return the host device discovered during runtime initialization.
-  static ol_device_handle_t getHostDevice();
+  ol_device_handle_t getHostDevice();
 
   /// Return the shared context that owns the discovered non-host devices.
-  static ol_context_handle_t getContext();
+  ol_context_handle_t getContext();
 
   /// Return the number of non-host devices available to kernel languages.
-  static int getDeviceCount();
+  int getDeviceCount();
 
   /// Register \p Kernel for the host-side kernel identifier \p ID.
   ///
   /// \p ID is the opaque kernel key emitted by Clang in the offload entry
   /// table.  It is later passed to the launch entry point to recover the
   /// corresponding liboffload symbol handle.
-  static void registerKernel(const void *ID, ol_symbol_handle_t Kernel);
+  void registerKernel(const void *ID, ol_symbol_handle_t Kernel);
 
   /// Remove any registered kernel handle for the host-side kernel key \p ID.
-  static void unregisterKernel(const void *ID);
+  void unregisterKernel(const void *ID);
 
   /// Return the registered kernel handle for the host-side kernel key \p ID.
-  static ol_symbol_handle_t getKernel(const void *ID);
+  ol_symbol_handle_t getKernel(const void *ID);
 
   /// Register \p Program for the binary image identifier \p ID.
   ///
   /// \p ID is the device image start address from the offload binary
   /// descriptor.  It keys the loaded program so later function registration
   /// can look up the program that owns each kernel symbol.
-  static void registerProgram(const void *ID, ol_program_handle_t Program);
+  void registerProgram(const void *ID, ol_program_handle_t Program);
 
   /// Remove and return the loaded program handle for binary image key \p ID.
-  static ol_program_handle_t unregisterProgram(const void *ID);
+  ol_program_handle_t unregisterProgram(const void *ID);
 
   /// Return the loaded program handle for binary image key \p ID.
-  static ol_program_handle_t getProgram(const void *ID);
+  ol_program_handle_t getProgram(const void *ID);
+
+  /// Return all streams currently known for \p Device.
+  SmallPtrSet<StreamTy *, 8> getDeviceStreams(ol_device_handle_t Device);
+
+  /// Return all explicitly created blocking streams for \p Device.
+  SmallPtrSet<StreamTy *, 8> getBlockingStreams(ol_device_handle_t Device);
+
+  /// Return true if \p Device has an existing legacy default stream.
+  bool hasLegacyDefaultStream(ol_device_handle_t Device);
+
+  /// Create a stream for \p Device and register it with the process state.
+  ol_result_t createStream(ol_device_handle_t Device, QueueKind Kind,
+                           StreamTy **Stream);
+
+  /// Destroy \p Stream after removing it from the process state.
+  ol_result_t destroyStream(StreamTy *Stream);
+
+  /// Return true if \p Stream is currently registered with the process state.
+  bool isStreamRegistered(StreamTy *Stream);
 
 private:
-  static StateTy &get();
-  static StateTy *tryGet();
   static bool addDevices(ol_device_handle_t Device, void *Payload);
 
-  llvm::ArrayRef<ol_device_handle_t> getDevices() const;
+  ArrayRef<ol_device_handle_t> getDevices() const;
 
   void addDevice(ol_device_handle_t Device);
   void setHostDevice(ol_device_handle_t Device);
 
-  void addKernel(KernelIDTy KernelID, ol_symbol_handle_t Kernel);
-  void removeKernel(KernelIDTy KernelID);
-  ol_symbol_handle_t lookupKernel(KernelIDTy KernelID);
+  StreamTy *getOrCreateDefaultStream(ol_device_handle_t Device);
+  void destroyDefaultStreams();
 
-  void addProgram(const void *Binary, ol_program_handle_t Program);
-  ol_program_handle_t removeProgram(const void *Binary);
-  ol_program_handle_t lookupProgram(const void *Binary);
+  /// Inserts the Stream into the DeviceStreamsMap and the
+  /// DeviceBlockingStreamsMap
+  void addStream(StreamTy *Stream);
+  /// Removes the Stream from DeviceStreamsMap and DeviceBlockingStreamsMap
+  void removeStream(StreamTy *Stream);
 
+  void destroyRegisteredStreams();
   void destroyRegisteredPrograms();
 
-  llvm::DenseMap<const void *, ol_program_handle_t> BinaryRegisterMap;
-  llvm::DenseMap<KernelIDTy, ol_symbol_handle_t> KernelMap;
-  llvm::SmallVector<ol_device_handle_t, 8> Devices;
+  DenseMap<const void *, ol_program_handle_t> BinaryRegisterMap;
+  DenseMap<KernelIDTy, ol_symbol_handle_t> KernelMap;
+  SmallVector<ol_device_handle_t, AssumedDeviceCount> Devices;
+  DenseMap<ol_device_handle_t, StreamTy *> DeviceDefaultStreamsMap;
+  DenseMap<ol_device_handle_t, SmallPtrSet<StreamTy *, AssumedStreamCount>>
+      DeviceStreamsMap;
+  DenseMap<ol_device_handle_t, SmallPtrSet<StreamTy *, AssumedStreamCount>>
+      DeviceBlockingStreamsMap;
 
   ol_context_handle_t Context = nullptr;
-  ol_queue_handle_t DefaultQueue = nullptr;
   ol_device_handle_t HostDevice = nullptr;
 
+  std::mutex DeviceDefaultStreamsMapLock;
+  std::mutex DeviceStreamsMapLock;
+  std::mutex DeviceBlockingStreamsMapLock;
+
   StateTy();
 };
 
diff --git a/offload/languages/kernel/include/Stream.h b/offload/languages/kernel/include/Stream.h
new file mode 100644
index 0000000000000..e66fc7000281d
--- /dev/null
+++ b/offload/languages/kernel/include/Stream.h
@@ -0,0 +1,33 @@
+//===-- Stream.h - Kernel language stream state ---------------------------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+
+#ifndef LLVM_OFFLOAD_LANGUAGES_KERNEL_INCLUDE_STREAM_H
+#define LLVM_OFFLOAD_LANGUAGES_KERNEL_INCLUDE_STREAM_H
+
+#include "OffloadAPI.h"
+
+namespace llvm {
+namespace offload {
+
+enum class QueueKind {
+  LegacyDefault,
+  PerThreadDefault,
+  ExplicitBlocking,
+  ExplicitNonBlocking,
+};
+
+struct StreamTy {
+  ol_queue_handle_t Queue = nullptr;
+  ol_device_handle_t Device = nullptr;
+  QueueKind Kind = QueueKind::ExplicitBlocking;
+};
+
+} // namespace offload
+} // namespace llvm
+
+#endif // LLVM_OFFLOAD_LANGUAGES_KERNEL_INCLUDE_STREAM_H
diff --git a/offload/languages/kernel/src/LanguageErrors.cpp b/offload/languages/kernel/src/LanguageErrors.cpp
index 823b2e34e2124..300271975cf94 100644
--- a/offload/languages/kernel/src/LanguageErrors.cpp
+++ b/offload/languages/kernel/src/LanguageErrors.cpp
@@ -58,13 +58,14 @@ const char *GetErrorString(Error_t Error) {
 }
 
 Error_t GetLastError() {
-  Error_t Error = static_cast<Error_t>(ThreadState::getLastError());
-  ThreadState::setLastError(Success);
+  ThreadState &State = ThreadState::get();
+  Error_t Error = static_cast<Error_t>(State.getLastError());
+  State.setLastError(Success);
   return Error;
 }
 
 Error_t PeekAtLastError() {
-  return static_cast<Error_t>(ThreadState::getLastError());
+  return static_cast<Error_t>(ThreadState::get().getLastError());
 }
 
 #include "UndefineLanguageNames.inc"
diff --git a/offload/languages/kernel/src/LanguageLaunch.cpp b/offload/languages/kernel/src/LanguageLaunch.cpp
index 7199e3a3f10f0..edc1209516ad6 100644
--- a/offload/languages/kernel/src/LanguageLaunch.cpp
+++ b/offload/languages/kernel/src/LanguageLaunch.cpp
@@ -8,30 +8,21 @@
 
 #include "LanguageLaunch.h"
 #include "LanguageUtils.h"
+#include "OffloadErrors.h"
 #include "State.h"
+#include "Stream.h"
 #include <cstdio>
 
-using RuntimeState = llvm::offload::StateTy;
-using ThreadState = llvm::offload::ThreadStateTy;
-
-static constexpr ol_error_struct_t InvalidKernelError = {
-    OL_ERRC_INVALID_NULL_HANDLE, "kernel is not registered"};
-
-static constexpr ol_error_struct_t InvalidDeviceError = {OL_ERRC_INVALID_DEVICE,
-                                                         "invalid device"};
-
-static constexpr ol_error_struct_t InvalidArgumentError = {
-    OL_ERRC_INVALID_ARGUMENT, "invalid argument to kernel launch"};
-
-static constexpr ol_error_struct_t InvalidConfigurationError = {
-    OL_ERRC_INVALID_SIZE, "invalid kernel launch configuration"};
+using namespace llvm::offload;
 
 /// Internal kernel launch implementation
 ol_result_t __llvmLaunchKernelImpl(const char *KernelID, dim3 GridDim,
                                    dim3 BlockDim, void *KernelArgsPtr,
                                    size_t DynamicSharedMem, void *Stream) {
-  ol_device_handle_t Device = ThreadState::getDefaultDevice();
-  ol_symbol_handle_t Kernel = RuntimeState::getKernel(KernelID);
+  StateTy &State = StateTy::get();
+  ThreadStateTy &ThreadState = ThreadStateTy::get();
+  ol_device_handle_t Device = ThreadState.getDefaultDevice();
+  ol_symbol_handle_t Kernel = State.getKernel(KernelID);
   if (!Device)
     return &InvalidDeviceError;
   if (!KernelID || !KernelArgsPtr)
@@ -54,8 +45,8 @@ ol_result_t __llvmLaunchKernelImpl(const char *KernelID, dim3 GridDim,
   LaunchSizeArgs.GroupSize.z = BlockDim.z;
   LaunchSizeArgs.DynSharedMemory = DynamicSharedMem;
 
-  ol_queue_handle_t Queue = Stream ? reinterpret_cast<ol_queue_handle_t>(Stream)
-                                   : ThreadState::getDefaultQueue();
+  ol_queue_handle_t Queue = Stream ? reinterpret_cast<StreamTy *>(Stream)->Queue
+                                   : ThreadState.getDefaultQueue();
 
   struct OffloadKernelArgs {
     void **Args;
@@ -81,7 +72,8 @@ extern "C" {
 /// Push call configuration for kernel launch
 unsigned __llvmPushCallConfiguration(dim3 GridSize, dim3 BlockSize,
                                      size_t SharedMemory, void *Stream) {
-  CallConfigurationTy &CC = ThreadState::getCallConfiguration();
+  ThreadStateTy &ThreadState = ThreadStateTy::get();
+  CallConfigurationTy &CC = ThreadState.getCallConfiguration();
 
   CC.GridSize = GridSize;
   CC.BlockSize = BlockSize;
@@ -93,7 +85,8 @@ unsigned __llvmPushCallConfiguration(dim3 GridSize, dim3 BlockSize,
 /// Pop call configuration for kernel launch
 unsigned __llvmPopCallConfiguration(dim3 *GridSize, dim3 *BlockSize,
                                     size_t *SharedMemory, void **Stream) {
-  CallConfigurationTy &CC = ThreadState::getCallConfiguration();
+  ThreadStateTy &ThreadState = ThreadStateTy::get();
+  CallConfigurationTy &CC = ThreadState.getCallConfiguration();
   *GridSize = CC.GridSize;
   *BlockSize = CC.BlockSize;
   *SharedMemory = CC.SharedMemory;
diff --git a/offload/languages/kernel/src/LanguageRegistration.cpp b/offload/languages/kernel/src/LanguageRegistration.cpp
index 2219e16ea0365..77aecdff21f23 100644
--- a/offload/languages/kernel/src/LanguageRegistration.cpp
+++ b/offload/languages/kernel/src/LanguageRegistration.cpp
@@ -27,12 +27,13 @@ extern "C" {
 void __llvmRegisterFunction(const char *Binary, const char *KernelID,
                             char *KernelName, const char *KernelName1, int,
                             uint3 *, uint3 *, dim3 *, dim3 *, int *) {
+  RuntimeState &State = RuntimeState::get();
   ol_symbol_handle_t Kernel;
-  ol_program_handle_t Program = RuntimeState::getProgram(Binary);
+  ol_program_handle_t Program = State.getProgram(Binary);
   ol_result_t Result = olGetSymbol(
       Program, KernelName, ol_symbol_kind_t::OL_SYMBOL_KIND_KERNEL, &Kernel);
   CHECK_FATAL(Result, "Failed to get kernel symbol for " << KernelName);
-  RuntimeState::registerKernel(KernelID, Kernel);
+  State.registerKernel(KernelID, Kernel);
 }
 
 void __llvmRegisterVar(void **, char *, char *, const char *, int, int, int,
@@ -76,7 +77,9 @@ struct __tgt_bin_desc {
 
 void __tgt_register_lib(__tgt_bin_desc *Desc) {
   // TODO: For each device, lazily.
-  ol_device_handle_t Device = ThreadState::getDefaultDevice();
+  RuntimeState &State = RuntimeState::get();
+  ThreadState &Thread = ThreadState::get();
+  ol_device_handle_t Device = Thread.getDefaultDevice();
 
   for (int32_t I = 0, E = Desc->NumDeviceImages; I < E; ++I) {
     ol_program_handle_t Program = nullptr;
@@ -94,7 +97,7 @@ void __tgt_register_lib(__tgt_bin_desc *Desc) {
       abort();
     }
 
-    RuntimeState::registerProgram(DeviceImage.ImageStart, Program);
+    State.registerProgram(DeviceImage.ImageStart, Program);
 
     for (auto *Entry = DeviceImage.EntriesBegin;
          Entry != DeviceImage.EntriesEnd; ++Entry) {
@@ -108,16 +111,20 @@ void __tgt_register_lib(__tgt_bin_desc *Desc) {
 }
 
 void __tgt_unregister_lib(__tgt_bin_desc *Desc) {
+  RuntimeState *State = RuntimeState::tryGet();
+  if (!State)
+    return;
+
   for (int32_t I = 0, E = Desc->NumDeviceImages; I < E; ++I) {
     __tgt_device_image &DeviceImage = Desc->DeviceImages[I];
     for (auto *Entry = DeviceImage.EntriesBegin;
          Entry != DeviceImage.EntriesEnd; ++Entry) {
       if (!Entry->Size && !Entry->Flags)
-        RuntimeState::unregisterKernel((const char *)Entry->Address);
+        State->unregisterKernel((const char *)Entry->Address);
     }
 
     if (ol_program_handle_t Program =
-            RuntimeState::unregisterProgram(DeviceImage.ImageStart))
+            State->unregisterProgram(DeviceImage.ImageStart))
       olDestroyProgram(Program);
   }
 }
diff --git a/offload/languages/kernel/src/LanguageRuntime.cpp b/offload/languages/kernel/src/LanguageRuntime.cpp
index d294412d37e8f..bf1b135aa7d7c 100644
--- a/offload/languages/kernel/src/LanguageRuntime.cpp
+++ b/offload/languages/kernel/src/LanguageRuntime.cpp
@@ -19,6 +19,7 @@
 
 #include "LanguageUtils.h"
 #include "State.h"
+#include "Stream.h"
 #include "Types.h"
 
 #include "OffloadAPI.h"
@@ -28,11 +29,11 @@
 #include <cstdlib>
 #include <cstring>
 
-using RuntimeState = llvm::offload::StateTy;
-using ThreadState = llvm::offload::ThreadStateTy;
+using namespace llvm::offload;
 
 Error_t Malloc(void **DevPtr, size_t Size) {
-  ol_device_handle_t Device = ThreadState::getDefaultDevice();
+  ThreadStateTy &ThreadState = ThreadStateTy::get();
+  ol_device_handle_t Device = ThreadState.getDefaultDevice();
   ol_result_t Result = olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, DevPtr);
   return convertAndSetLastError(Result);
 }
@@ -43,30 +44,32 @@ Error_t Free(void *DevPtr) {
 }
 
 Error_t Memcpy(void *Dst, const void *Src, size_t Size, MemcpyKind Kind) {
-  ol_queue_handle_t Queue = ThreadState::getDefaultQueue();
+  StateTy &State = StateTy::get();
+  ThreadStateTy &ThreadState = ThreadStateTy::get();
+  ol_queue_handle_t Queue = ThreadState.getDefaultQueue();
 
   ol_result_t Result;
   switch (Kind) {
   case MemcpyHostToHost: {
-    ol_device_handle_t Host = RuntimeState::getHostDevice();
+    ol_device_handle_t Host = State.getHostDevice();
     Result = olMemcpy(nullptr, Dst, Host, const_cast<void *>(Src), Host, Size);
     break;
   }
   case MemcpyHostToDevice: {
-    ol_device_handle_t Device = ThreadState::getDefaultDevice();
-    ol_device_handle_t Host = RuntimeState::getHostDevice();
+    ol_device_handle_t Device = ThreadState.getDefaultDevice();
+    ol_device_handle_t Host = State.getHostDevice();
     Result = olMemcpy(Queue, Dst, Device, const_cast<void *>(Src), Host, Size);
     break;
   }
   case MemcpyDeviceToHost: {
-    ol_device_handle_t Device = ThreadState::getDefaultDevice();
-    ol_device_handle_t Host = RuntimeState::getHostDevice();
+    ol_device_handle_t Device = ThreadState.getDefaultDevice();
+    ol_device_handle_t Host = State.getHostDevice();
 
     Result = olMemcpy(Queue, Dst, Host, const_cast<void *>(Src), Device, Size);
     break;
   }
   case MemcpyDeviceToDevice: {
-    ol_device_handle_t Device = ThreadState::getDefaultDevice();
+    ol_device_handle_t Device = ThreadState.getDefaultDevice();
 
     Result =
         olMemcpy(Queue, Dst, Device, const_cast<void *>(Src), Device, Size);
@@ -86,34 +89,38 @@ Error_t Memcpy(void *Dst, const void *Src, size_t Size, MemcpyKind Kind) {
 Error_t DeviceSynchronize() {
   // TODO: This is not correct. We likely want to pipe this through to the
   // plugins.
-  ol_queue_handle_t Queue = ThreadState::getDefaultQueue();
+  ThreadStateTy &ThreadState = ThreadStateTy::get();
+  ol_queue_handle_t Queue = ThreadState.getDefaultQueue();
   ol_result_t Result = olSyncQueue(Queue);
   return convertAndSetLastError(Result);
 }
 
 Error_t GetDevice(int *DeviceNo) {
-  ol_device_handle_t Device = ThreadState::getDevice(DeviceNo);
+  ThreadStateTy &ThreadState = ThreadStateTy::get();
+  ol_device_handle_t Device = ThreadState.getDevice(DeviceNo);
   if (!Device)
     return setLastError(ErrorInvalidDevice);
   return setLastError(Success);
 }
 
 Error_t GetDeviceCount(int *Count) {
-  *Count = RuntimeState::getDeviceCount();
+  *Count = StateTy::get().getDeviceCount();
   return setLastError(Success);
 }
 
 Error_t SetDevice(int DeviceNo) {
-  ol_device_handle_t Device = ThreadState::setDefaultDevice(DeviceNo);
+  ThreadStateTy &ThreadState = ThreadStateTy::get();
+  ol_device_handle_t Device = ThreadState.setDefaultDevice(DeviceNo);
   if (!Device)
     return setLastError(ErrorInvalidDevice);
-  assert(Device == ThreadState::getDefaultDevice() &&
+  assert(Device == ThreadState.getDefaultDevice() &&
          "Set Device is not Default Device");
   return setLastError(Success);
 }
 
 Error_t HostAlloc(void **Ptr, size_t Size, unsigned int Flags) {
-  ol_device_handle_t Device = ThreadState::getDefaultDevice();
+  ThreadStateTy &ThreadState = ThreadStateTy::get();
+  ol_device_handle_t Device = ThreadState.getDefaultDevice();
   ol_result_t Result = olMemAllocHost(Device, Size, Ptr);
   return convertAndSetLastError(Result);
 }
@@ -128,7 +135,8 @@ Error_t FreeHost(void *Ptr) {
 }
 
 Error_t GetDeviceProperties(DeviceProp_t *DeviceProp, int DeviceNo) {
-  ol_device_handle_t Device = ThreadState::getDefaultDevice();
+  ThreadStateTy &ThreadState = ThreadStateTy::get();
+  ol_device_handle_t Device = ThreadState.getDefaultDevice();
   size_t NameSize = 0;
   olGetDeviceInfoSize(Device, OL_DEVICE_INFO_NAME, &NameSize);
   assert(NameSize <= sizeof(DeviceProp->name) &&
@@ -144,20 +152,18 @@ Error_t GetDeviceProperties(DeviceProp_t *DeviceProp, int DeviceNo) {
 }
 
 Error_t StreamCreate(Stream_t *Stream) {
-  ol_queue_handle_t Queue;
-  ol_result_t Result = olCreateQueue(RuntimeState::getContext(),
-                                     ThreadState::getDefaultDevice(), &Queue);
+  StreamTy *StreamObj = nullptr;
+  StateTy &State = StateTy::get();
+  ThreadStateTy &ThreadState = ThreadStateTy::get();
+  ol_result_t Result = State.createStream(
+      ThreadState.getDefaultDevice(), QueueKind::ExplicitBlocking, &StreamObj);
   if (Result == OL_SUCCESS)
-    *Stream = reinterpret_cast<Stream_t>(Queue);
+    *Stream = toLanguageStream(StreamObj);
   return convertAndSetLastError(Result);
 }
 
 Error_t StreamDestroy(Stream_t Stream) {
-  ol_queue_handle_t Queue;
-  Error_t Err = getQueueFromStream(Stream, &Queue);
-  if (Err != Success)
-    return setLastError(Err);
-  ol_result_t Result = olDestroyQueue(Queue);
+  ol_result_t Result = StateTy::get().destroyStream(toInternalStream(Stream));
   return convertAndSetLastError(Result);
 }
 
diff --git a/offload/languages/kernel/src/State.cpp b/offload/languages/kernel/src/State.cpp
index baaceeb377bc7..f18bd1d57beb4 100644
--- a/offload/languages/kernel/src/State.cpp
+++ b/offload/languages/kernel/src/State.cpp
@@ -7,9 +7,11 @@
 //===----------------------------------------------------------------------===//
 
 #include "State.h"
+#include "Stream.h"
 #include "Types.h"
 
 #include "OffloadAPI.h"
+#include "OffloadErrors.h"
 #include "llvm/ADT/ArrayRef.h"
 #include "llvm/ADT/DenseMap.h"
 #include "llvm/ADT/SmallPtrSet.h"
@@ -17,6 +19,7 @@
 
 #include <atomic>
 #include <cassert>
+#include <cstddef>
 #include <cstdint>
 #include <cstdio>
 #include <mutex>
@@ -60,32 +63,52 @@ static void deleteThreadStates() {
 }
 
 static void deleteState() {
-  StateTy *ST = StatePtr.load();
-  StatePtr.store(nullptr);
+  StateTy *ST = StatePtr.load(std::memory_order_acquire);
+  if (!ST)
+    return;
   delete ST;
-  StatePtr = nullptr;
+  StatePtr.store(nullptr, std::memory_order_release);
 }
 
-static void destroyQueue(ol_queue_handle_t &Queue) {
-  if (!Queue)
+static void destroyStreamHandle(StreamTy *&Stream) {
+  if (!Stream)
     return;
 
-  olSyncQueue(Queue);
-  olDestroyQueue(Queue);
-  Queue = nullptr;
+  olSyncQueue(Stream->Queue);
+  olDestroyQueue(Stream->Queue);
+  delete Stream;
+  Stream = nullptr;
 }
 
 namespace llvm {
 namespace offload {
 
+static bool removeStreamFromMap(
+    DenseMap<ol_device_handle_t, SmallPtrSet<StreamTy *, 8>> &StreamsMap,
+    StreamTy *Stream) {
+  bool Removed = false;
+  SmallVector<ol_device_handle_t, 8> EmptyDevices;
+
+  for (auto &It : StreamsMap) {
+    Removed |= It.second.erase(Stream);
+    if (It.second.empty())
+      EmptyDevices.push_back(It.first);
+  }
+
+  for (ol_device_handle_t Device : EmptyDevices)
+    StreamsMap.erase(Device);
+
+  return Removed;
+}
+
 // ThreadStateTy implementation.
 
 ThreadStateTy::ThreadStateTy() {
-  if (PerThreadQueue) [[unlikely]]
-    createDefaultQueue(getDefaultDevice());
+  unsigned int NumDevices = StateTy::get().Devices.size();
+  PerThreadDeviceDefaultStreamMap.reserve(NumDevices);
   atexit(deleteThreadStates);
 }
-ThreadStateTy::~ThreadStateTy() { destroyQueue(DefaultQueue); }
+ThreadStateTy::~ThreadStateTy() { destroyDefaultStreams(); }
 
 ThreadStateTy &ThreadStateTy::get() {
   auto *&TS = ThreadState;
@@ -100,141 +123,257 @@ ThreadStateTy &ThreadStateTy::get() {
 }
 
 ol_device_handle_t ThreadStateTy::getDefaultDevice() {
-  int DD = ThreadStateTy::get().DefaultDevice;
-  return StateTy::get().getDevices()[DD];
+  ArrayRef<ol_device_handle_t> Devices = StateTy::get().getDevices();
+  int DD = DefaultDevice;
+  if (DD < 0 || DD >= static_cast<int>(Devices.size()))
+    return nullptr;
+  return Devices[DD];
 }
 
-ol_queue_handle_t ThreadStateTy::getDefaultQueue() {
+StreamTy *ThreadStateTy::getDefaultStream() {
+  ol_device_handle_t Device = getDefaultDevice();
+  if (!Device)
+    return nullptr;
+
   if (!PerThreadQueue) [[likely]]
-    return StateTy::get().DefaultQueue;
-  return ThreadStateTy::get().DefaultQueue;
+    return StateTy::get().getOrCreateDefaultStream(Device);
+
+  return getOrCreateDefaultStream(Device);
 }
 
-CallConfigurationTy &ThreadStateTy::getCallConfiguration() {
-  return ThreadStateTy::get().CC;
+ol_queue_handle_t ThreadStateTy::getDefaultQueue() {
+  if (StreamTy *Stream = getDefaultStream())
+    return Stream->Queue;
+  return nullptr;
 }
 
+CallConfigurationTy &ThreadStateTy::getCallConfiguration() { return CC; }
+
 ol_device_handle_t ThreadStateTy::setDefaultDevice(int DeviceNo) {
   ArrayRef<ol_device_handle_t> Devices = StateTy::get().getDevices();
   if (DeviceNo < 0 || DeviceNo >= static_cast<int>(Devices.size()))
     return nullptr;
-  ThreadStateTy::get().DefaultDevice = DeviceNo;
-  ol_device_handle_t DD = Devices[DeviceNo];
-  ThreadStateTy::get().createDefaultQueue(DD);
-  return DD;
+  DefaultDevice = DeviceNo;
+  return Devices[DeviceNo];
 }
 
 ol_device_handle_t ThreadStateTy::getDevice(int *DeviceNo) {
-  *DeviceNo = ThreadStateTy::get().DefaultDevice;
-  return ThreadStateTy::getDefaultDevice();
+  *DeviceNo = DefaultDevice;
+  return getDefaultDevice();
 }
 
-uint32_t ThreadStateTy::getLastError() {
-  return ThreadStateTy::get().LastError;
-}
+uint32_t ThreadStateTy::getLastError() { return LastError; }
 
 uint32_t ThreadStateTy::setLastError(uint32_t Error) {
-  return ThreadStateTy::get().LastError = Error;
+  return LastError = Error;
+}
+
+StreamTy *ThreadStateTy::getOrCreateDefaultStream(ol_device_handle_t Device) {
+  if (!Device)
+    return nullptr;
+
+  StateTy &State = StateTy::get();
+  ol_context_handle_t Context = State.getContext();
+  if (!Context)
+    return nullptr;
+
+  StreamTy *&Stream = PerThreadDeviceDefaultStreamMap[Device];
+  if (!Stream) {
+    ol_queue_handle_t Queue = nullptr;
+    CHECK_FATAL(olCreateQueue(Context, Device, &Queue),
+                "Failed to create per-thread default queue for device");
+    Stream = new StreamTy{Queue, Device, QueueKind::PerThreadDefault};
+    State.addStream(Stream);
+  }
+  return Stream;
 }
 
-void ThreadStateTy::createDefaultQueue(ol_device_handle_t Device) {
-  if (DefaultQueue)
-    olDestroyQueue(DefaultQueue);
-  CHECK_FATAL(olCreateQueue(StateTy::getContext(), Device, &DefaultQueue),
-              "Failed to create per-thread default queue");
+void ThreadStateTy::destroyDefaultStreams() {
+  for (auto &It : PerThreadDeviceDefaultStreamMap) {
+    if (StateTy *State = StateTy::tryGet())
+      State->removeStream(It.second);
+    destroyStreamHandle(It.second);
+  }
+  PerThreadDeviceDefaultStreamMap.clear();
 }
 
 // StateTy implementation.
 
 StateTy &StateTy::get() {
-  StateTy *ST = StatePtr.load();
+  StateTy *ST = StatePtr.load(std::memory_order_acquire);
   if (!ST) [[unlikely]] {
     std::lock_guard<std::mutex> LG(getStateLock());
-    ST = StatePtr.load();
+    ST = StatePtr.load(std::memory_order_acquire);
     if (!ST) {
       ST = new StateTy();
-      StatePtr.store(ST);
+      StatePtr.store(ST, std::memory_order_release);
     }
   }
   return *ST;
 }
 
-StateTy *StateTy::tryGet() { return StatePtr.load(); }
+StateTy *StateTy::tryGet() { return StatePtr.load(std::memory_order_acquire); }
 
-ol_device_handle_t StateTy::getHostDevice() { return get().HostDevice; }
+StreamTy *StateTy::getOrCreateDefaultStream(ol_device_handle_t Device) {
+  if (!Device)
+    return nullptr;
 
-ol_context_handle_t StateTy::getContext() { return get().Context; }
+  ol_context_handle_t RuntimeContext = getContext();
+  if (!RuntimeContext)
+    return nullptr;
 
-int StateTy::getDeviceCount() {
-  int DeviceCount = get().getDevices().size();
-  return DeviceCount;
+  std::lock_guard<std::mutex> LG(DeviceDefaultStreamsMapLock);
+  StreamTy *&Stream = DeviceDefaultStreamsMap[Device];
+  if (!Stream) {
+    ol_queue_handle_t Queue = nullptr;
+    CHECK_FATAL(olCreateQueue(RuntimeContext, Device, &Queue),
+                "Failed to create default queue for device");
+    Stream = new StreamTy{Queue, Device, QueueKind::LegacyDefault};
+    addStream(Stream);
+  }
+  return Stream;
 }
 
-ArrayRef<ol_device_handle_t> StateTy::getDevices() const { return Devices; }
+void StateTy::destroyDefaultStreams() {
+  std::lock_guard<std::mutex> LG(DeviceDefaultStreamsMapLock);
+  for (auto &It : DeviceDefaultStreamsMap) {
+    removeStream(It.second);
+    destroyStreamHandle(It.second);
+  }
+  DeviceDefaultStreamsMap.clear();
+}
 
-void StateTy::addDevice(ol_device_handle_t Device) {
-  Devices.push_back(Device);
+void StateTy::addStream(StreamTy *Stream) {
+  std::lock_guard<std::mutex> LG(DeviceStreamsMapLock);
+  DeviceStreamsMap[Stream->Device].insert(Stream);
 }
 
-void StateTy::setHostDevice(ol_device_handle_t Device) {
-  if (!HostDevice)
-    HostDevice = Device;
+void StateTy::removeStream(StreamTy *Stream) {
+  if (!Stream)
+    return;
+
+  {
+    std::lock_guard<std::mutex> LG(DeviceStreamsMapLock);
+    if (!removeStreamFromMap(DeviceStreamsMap, Stream))
+      return;
+  }
+
+  std::lock_guard<std::mutex> LG(DeviceBlockingStreamsMapLock);
+  removeStreamFromMap(DeviceBlockingStreamsMap, Stream);
 }
 
-void StateTy::addKernel(KernelIDTy KernelID, ol_symbol_handle_t Kernel) {
-  KernelMap[KernelID] = Kernel;
+SmallPtrSet<StreamTy *, 8>
+StateTy::getDeviceStreams(ol_device_handle_t Device) {
+  std::lock_guard<std::mutex> LG(DeviceStreamsMapLock);
+  auto It = DeviceStreamsMap.find(Device);
+  if (It == DeviceStreamsMap.end())
+    return {};
+  return It->second;
 }
 
-void StateTy::removeKernel(KernelIDTy KernelID) { KernelMap.erase(KernelID); }
+SmallPtrSet<StreamTy *, 8>
+StateTy::getBlockingStreams(ol_device_handle_t Device) {
+  std::lock_guard<std::mutex> LG(DeviceBlockingStreamsMapLock);
+  auto It = DeviceBlockingStreamsMap.find(Device);
+  if (It == DeviceBlockingStreamsMap.end())
+    return {};
+  return It->second;
+}
 
-ol_symbol_handle_t StateTy::lookupKernel(KernelIDTy KernelID) {
-  return KernelMap[KernelID];
+bool StateTy::hasLegacyDefaultStream(ol_device_handle_t Device) {
+  std::lock_guard<std::mutex> LG(DeviceDefaultStreamsMapLock);
+  auto It = DeviceDefaultStreamsMap.find(Device);
+  return It != DeviceDefaultStreamsMap.end() && It->second;
 }
 
-void StateTy::registerKernel(const void *ID, ol_symbol_handle_t Kernel) {
-  get().addKernel(ID, Kernel);
+ol_result_t StateTy::createStream(ol_device_handle_t Device, QueueKind Kind,
+                                  StreamTy **Stream) {
+  if (!Stream)
+    return &InvalidNullPointerError;
+  *Stream = nullptr;
+
+  ol_context_handle_t RuntimeContext = getContext();
+  if (!Device || !RuntimeContext)
+    return &InvalidDeviceError;
+
+  ol_queue_handle_t Queue = nullptr;
+  ol_result_t Result = olCreateQueue(RuntimeContext, Device, &Queue);
+  if (Result == OL_SUCCESS) {
+    *Stream = new StreamTy{Queue, Device, Kind};
+    addStream(*Stream);
+    if (Kind == QueueKind::ExplicitBlocking) {
+      std::lock_guard<std::mutex> LG(DeviceBlockingStreamsMapLock);
+      DeviceBlockingStreamsMap[Device].insert(*Stream);
+    }
+  }
+  return Result;
 }
 
-void StateTy::unregisterKernel(const void *ID) {
-  if (StateTy *State = tryGet())
-    State->removeKernel(ID);
+ol_result_t StateTy::destroyStream(StreamTy *Stream) {
+  if (!Stream)
+    return &InvalidNullPointerError;
+
+  if (!isStreamRegistered(Stream))
+    return &InvalidStreamError;
+
+  removeStream(Stream);
+  ol_result_t Result = olDestroyQueue(Stream->Queue);
+  delete Stream;
+  return Result;
 }
 
-ol_symbol_handle_t StateTy::getKernel(const void *ID) {
-  return get().lookupKernel(ID);
+bool StateTy::isStreamRegistered(StreamTy *Stream) {
+  if (!Stream)
+    return false;
+
+  std::lock_guard<std::mutex> LG(DeviceStreamsMapLock);
+  for (auto &It : DeviceStreamsMap)
+    if (It.second.contains(Stream))
+      return true;
+  return false;
 }
 
-void StateTy::addProgram(const void *Binary, ol_program_handle_t Program) {
-  BinaryRegisterMap[Binary] = Program;
+ol_device_handle_t StateTy::getHostDevice() { return HostDevice; }
+
+ol_context_handle_t StateTy::getContext() { return Context; }
+
+int StateTy::getDeviceCount() { return Devices.size(); }
+
+ArrayRef<ol_device_handle_t> StateTy::getDevices() const { return Devices; }
+
+void StateTy::addDevice(ol_device_handle_t Device) {
+  Devices.push_back(Device);
 }
 
-ol_program_handle_t StateTy::removeProgram(const void *Binary) {
-  auto It = BinaryRegisterMap.find(Binary);
-  if (It == BinaryRegisterMap.end())
-    return nullptr;
-  ol_program_handle_t Program = It->second;
-  BinaryRegisterMap.erase(It);
-  return Program;
+void StateTy::setHostDevice(ol_device_handle_t Device) {
+  if (!HostDevice)
+    HostDevice = Device;
 }
 
-ol_program_handle_t StateTy::lookupProgram(const void *Binary) {
-  assert(BinaryRegisterMap.count(Binary) &&
-         "Program not registered for binary");
-  return BinaryRegisterMap[Binary];
+void StateTy::registerKernel(const void *ID, ol_symbol_handle_t Kernel) {
+  KernelMap[ID] = Kernel;
 }
 
+void StateTy::unregisterKernel(const void *ID) { KernelMap.erase(ID); }
+
+ol_symbol_handle_t StateTy::getKernel(const void *ID) { return KernelMap[ID]; }
+
 void StateTy::registerProgram(const void *ID, ol_program_handle_t Program) {
-  get().addProgram(ID, Program);
+  BinaryRegisterMap[ID] = Program;
 }
 
 ol_program_handle_t StateTy::unregisterProgram(const void *ID) {
-  if (StateTy *State = tryGet())
-    return State->removeProgram(ID);
-  return nullptr;
+  auto It = BinaryRegisterMap.find(ID);
+  if (It == BinaryRegisterMap.end())
+    return nullptr;
+  ol_program_handle_t Program = It->second;
+  BinaryRegisterMap.erase(It);
+  return Program;
 }
 
 ol_program_handle_t StateTy::getProgram(const void *ID) {
-  return get().lookupProgram(ID);
+  assert(BinaryRegisterMap.count(ID) && "Program not registered for binary");
+  return BinaryRegisterMap[ID];
 }
 
 bool StateTy::addDevices(ol_device_handle_t Device, void *Payload) {
@@ -269,23 +408,41 @@ StateTy::StateTy() {
     CHECK_FATAL(olCreateContext(Devices.size(), Devices.data(), &Context),
                 "Failed to create default context");
 
-  if (!PerThreadQueue) [[likely]]
-    if (!Devices.empty()) [[likely]]
-      CHECK_FATAL(olCreateQueue(Context, Devices.front(), &DefaultQueue),
-                  "Failed to create default queue");
+  unsigned int DeviceCount = Devices.size();
+  DeviceDefaultStreamsMap.reserve(DeviceCount);
+  DeviceStreamsMap.reserve(DeviceCount);
+  DeviceBlockingStreamsMap.reserve(DeviceCount);
 
   atexit(deleteState);
 }
 
 StateTy::~StateTy() {
   deleteThreadStates();
-  destroyQueue(DefaultQueue);
+  destroyDefaultStreams();
+  destroyRegisteredStreams();
   destroyRegisteredPrograms();
   if (Context)
     olDestroyContext(Context);
   olShutDown();
 }
 
+void StateTy::destroyRegisteredStreams() {
+  SmallVector<StreamTy *, 16> Streams;
+  {
+    std::lock_guard<std::mutex> LG(DeviceStreamsMapLock);
+    for (auto &It : DeviceStreamsMap)
+      Streams.append(It.second.begin(), It.second.end());
+    DeviceStreamsMap.clear();
+  }
+  {
+    std::lock_guard<std::mutex> LG(DeviceBlockingStreamsMapLock);
+    DeviceBlockingStreamsMap.clear();
+  }
+
+  for (StreamTy *&Stream : Streams)
+    destroyStreamHandle(Stream);
+}
+
 void StateTy::destroyRegisteredPrograms() {
   SmallPtrSet<ol_program_handle_t, 8> Programs;
   for (auto &It : BinaryRegisterMap)
diff --git a/offload/test/offloading/CUDA/stream_api.cu b/offload/test/offloading/CUDA/stream_api.cu
index 7202751f8207e..4a659b5a106f4 100644
--- a/offload/test/offloading/CUDA/stream_api.cu
+++ b/offload/test/offloading/CUDA/stream_api.cu
@@ -14,7 +14,12 @@
 
 #include <stdio.h>
 
-__global__ void setValue(int *Out) { *Out = 42; }
+static void print_error(const char *Label, cudaError_t Error) {
+  printf("%s value: %u\n", Label, static_cast<unsigned>(Error));
+  printf("%s name: %s\n", Label, cudaGetErrorName(Error));
+}
+
+__global__ void setValue(int *Out, int Value) { *Out = Value; }
 
 int main(int argc, char **argv) {
   cudaStream_t Stream = nullptr;
@@ -24,23 +29,40 @@ int main(int argc, char **argv) {
   printf("stream created: %d\n", Stream != nullptr);
   // CHECK: stream created: 1
 
-  int *DevPtr = nullptr;
-  int Result = 0;
-  if (cudaMalloc(&DevPtr, sizeof(int)) != cudaSuccess)
+  int *StreamPtr = nullptr;
+  int *DefaultPtr = nullptr;
+  int StreamResult = 0;
+  int DefaultResult = 0;
+  if (cudaMalloc(&StreamPtr, sizeof(int)) != cudaSuccess)
+    return 1;
+  if (cudaMalloc(&DefaultPtr, sizeof(int)) != cudaSuccess)
     return 1;
 
-  setValue<<<1, 1, 0, Stream>>>(DevPtr);
+  setValue<<<1, 1, 0, Stream>>>(StreamPtr, 42);
+  setValue<<<1, 1>>>(DefaultPtr, 17);
 
   if (cudaStreamSynchronize(Stream) != cudaSuccess)
     return 1;
-  if (cudaMemcpy(&Result, DevPtr, sizeof(int), cudaMemcpyDeviceToHost) !=
-      cudaSuccess)
+  if (cudaDeviceSynchronize() != cudaSuccess)
+    return 1;
+  if (cudaMemcpy(&StreamResult, StreamPtr, sizeof(int),
+                 cudaMemcpyDeviceToHost) != cudaSuccess)
+    return 1;
+  if (cudaMemcpy(&DefaultResult, DefaultPtr, sizeof(int),
+                 cudaMemcpyDeviceToHost) != cudaSuccess)
     return 1;
 
-  printf("stream result: %d\n", Result);
+  printf("stream result: %d\n", StreamResult);
   // CHECK: stream result: 42
+  printf("default result: %d\n", DefaultResult);
+  // CHECK: default result: 17
 
   if (cudaStreamDestroy(Stream) != cudaSuccess)
     return 1;
-  cudaFree(DevPtr);
+  print_error("destroyed stream destroy", cudaStreamDestroy(Stream));
+  // CHECK: destroyed stream destroy value: 4
+  // CHECK: destroyed stream destroy name: cudaErrorInvalidResourceHandle
+
+  cudaFree(StreamPtr);
+  cudaFree(DefaultPtr);
 }
diff --git a/offload/test/offloading/HIP/stream_api.hip b/offload/test/offloading/HIP/stream_api.hip
index c0e2699822814..0f5c6fd63ab81 100644
--- a/offload/test/offloading/HIP/stream_api.hip
+++ b/offload/test/offloading/HIP/stream_api.hip
@@ -14,7 +14,12 @@
 
 #include <stdio.h>
 
-__global__ void setValue(int *Out) { *Out = 42; }
+static void print_error(const char *Label, hipError_t Error) {
+  printf("%s value: %u\n", Label, static_cast<unsigned>(Error));
+  printf("%s name: %s\n", Label, hipGetErrorName(Error));
+}
+
+__global__ void setValue(int *Out, int Value) { *Out = Value; }
 
 int main(int argc, char **argv) {
   hipStream_t Stream = nullptr;
@@ -24,23 +29,40 @@ int main(int argc, char **argv) {
   printf("stream created: %d\n", Stream != nullptr);
   // CHECK: stream created: 1
 
-  int *DevPtr = nullptr;
-  int Result = 0;
-  if (hipMalloc(&DevPtr, sizeof(int)) != hipSuccess)
+  int *StreamPtr = nullptr;
+  int *DefaultPtr = nullptr;
+  int StreamResult = 0;
+  int DefaultResult = 0;
+  if (hipMalloc(&StreamPtr, sizeof(int)) != hipSuccess)
+    return 1;
+  if (hipMalloc(&DefaultPtr, sizeof(int)) != hipSuccess)
     return 1;
 
-  setValue<<<1, 1, 0, Stream>>>(DevPtr);
+  setValue<<<1, 1, 0, Stream>>>(StreamPtr, 42);
+  setValue<<<1, 1>>>(DefaultPtr, 17);
 
   if (hipStreamSynchronize(Stream) != hipSuccess)
     return 1;
-  if (hipMemcpy(&Result, DevPtr, sizeof(int), hipMemcpyDeviceToHost) !=
+  if (hipDeviceSynchronize() != hipSuccess)
+    return 1;
+  if (hipMemcpy(&StreamResult, StreamPtr, sizeof(int), hipMemcpyDeviceToHost) !=
       hipSuccess)
     return 1;
+  if (hipMemcpy(&DefaultResult, DefaultPtr, sizeof(int),
+                hipMemcpyDeviceToHost) != hipSuccess)
+    return 1;
 
-  printf("stream result: %d\n", Result);
+  printf("stream result: %d\n", StreamResult);
   // CHECK: stream result: 42
+  printf("default result: %d\n", DefaultResult);
+  // CHECK: default result: 17
 
   if (hipStreamDestroy(Stream) != hipSuccess)
     return 1;
-  hipFree(DevPtr);
+  print_error("destroyed stream destroy", hipStreamDestroy(Stream));
+  // CHECK: destroyed stream destroy value: 4
+  // CHECK: destroyed stream destroy name: hipErrorInvalidResourceHandle
+
+  hipFree(StreamPtr);
+  hipFree(DefaultPtr);
 }



More information about the llvm-commits mailing list