[llvm-branch-commits] [llvm] [LLVMOffload] Add StreamCreateWithFlags (PR #216381)

Sophia Herrmann via llvm-branch-commits llvm-branch-commits at lists.llvm.org
Mon Aug 17 11:14:38 PDT 2026


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

>From c134e1c95ad7e4dca977526326aa193d02ff64f7 Mon Sep 17 00:00:00 2001
From: Sophia Herrmann <herrmann15 at llnl.gov>
Date: Thu, 13 Aug 2026 16:38:44 -0700
Subject: [PATCH] add StreamCreateWithFlags

---
 .../include/kernel/DefineLanguageNames.inc    |  3 +++
 .../include/kernel/LanguageRuntime.h          | 11 +++++++-
 .../include/kernel/UndefineLanguageNames.inc  |  4 ++-
 .../languages/kernel/src/LanguageRuntime.cpp  | 24 +++++++++++++++++
 offload/test/offloading/CUDA/stream_api.cu    | 27 +++++++++++++++++++
 offload/test/offloading/HIP/stream_api.hip    | 26 ++++++++++++++++++
 6 files changed, 93 insertions(+), 2 deletions(-)

diff --git a/offload/languages/include/kernel/DefineLanguageNames.inc b/offload/languages/include/kernel/DefineLanguageNames.inc
index 790fd3ad1a359..f4732002c1b52 100644
--- a/offload/languages/include/kernel/DefineLanguageNames.inc
+++ b/offload/languages/include/kernel/DefineLanguageNames.inc
@@ -43,6 +43,9 @@
 #define FreeHost COMBINE(LANGUAGE, FreeHost)
 #define GetDeviceProperties COMBINE(LANGUAGE, GetDeviceProperties)
 #define Stream_t COMBINE(LANGUAGE, Stream_t)
+#define StreamDefault COMBINE(LANGUAGE, StreamDefault)
+#define StreamNonBlocking COMBINE(LANGUAGE, StreamNonBlocking)
 #define StreamCreate COMBINE(LANGUAGE, StreamCreate)
+#define StreamCreateWithFlags COMBINE(LANGUAGE, StreamCreateWithFlags)
 #define StreamDestroy COMBINE(LANGUAGE, StreamDestroy)
 #define StreamSynchronize COMBINE(LANGUAGE, StreamSynchronize)
diff --git a/offload/languages/include/kernel/LanguageRuntime.h b/offload/languages/include/kernel/LanguageRuntime.h
index 3fbc4c0af3e2c..b25e735004ee9 100644
--- a/offload/languages/include/kernel/LanguageRuntime.h
+++ b/offload/languages/include/kernel/LanguageRuntime.h
@@ -40,13 +40,20 @@ enum MemcpyKind {
   MemcpyDefault = 4
 };
 
-enum HostAllocFlags : unsigned int {
+/// Flags passed to HostAlloc
+enum : unsigned int {
   HostAllocDefault = 0x00,
   HostAllocPortable = 0x01,
   HostAllocMapped = 0x02,
   HostAllocWriteCombined = 0x04,
 };
 
+/// Flags passed to StreamCreateWithFlags
+enum : unsigned int {
+  StreamDefault = 0x00,
+  StreamNonBlocking = 0x01,
+};
+
 typedef struct Stream_st *Stream_t;
 
 /// Malloc, with type template overlay.
@@ -99,6 +106,8 @@ Error_t GetDeviceProperties(DeviceProp_t *DeviceProp, int DeviceNo);
 
 Error_t StreamCreate(Stream_t *stream);
 
+Error_t StreamCreateWithFlags(Stream_t *stream, unsigned int flags);
+
 Error_t StreamDestroy(Stream_t stream);
 
 Error_t StreamSynchronize(Stream_t stream);
diff --git a/offload/languages/include/kernel/UndefineLanguageNames.inc b/offload/languages/include/kernel/UndefineLanguageNames.inc
index cd0e6b33aafc6..0ec241ead0434 100644
--- a/offload/languages/include/kernel/UndefineLanguageNames.inc
+++ b/offload/languages/include/kernel/UndefineLanguageNames.inc
@@ -32,7 +32,6 @@
 #undef GetDeviceCount
 #undef SetDevice
 #undef HostAlloc
-#undef HostAllocFlags
 #undef HostAllocDefault
 #undef HostAllocPortable
 #undef HostAllocMapped
@@ -41,6 +40,9 @@
 #undef FreeHost
 #undef GetDeviceProperties
 #undef Stream_t
+#undef StreamDefault
+#undef StreamNonBlocking
 #undef StreamCreate
+#undef StreamCreateWithFlags
 #undef StreamDestroy
 #undef StreamSynchronize
diff --git a/offload/languages/kernel/src/LanguageRuntime.cpp b/offload/languages/kernel/src/LanguageRuntime.cpp
index c4a2cd6d18cb9..f0298302b2889 100644
--- a/offload/languages/kernel/src/LanguageRuntime.cpp
+++ b/offload/languages/kernel/src/LanguageRuntime.cpp
@@ -31,6 +31,7 @@
 
 using RuntimeState = llvm::offload::StateTy;
 using ThreadState = llvm::offload::ThreadStateTy;
+using StreamTy = llvm::offload::StreamTy;
 
 Error_t Malloc(void **DevPtr, size_t Size) {
   ol_device_handle_t Device = ThreadState::getDefaultDevice();
@@ -145,6 +146,9 @@ Error_t GetDeviceProperties(DeviceProp_t *DeviceProp, int DeviceNo) {
 }
 
 Error_t StreamCreate(Stream_t *Stream) {
+  if (!Stream)
+    return setLastError(ErrorInvalidValue);
+
   llvm::offload::StreamTy *StreamObj = nullptr;
   ol_result_t Result = RuntimeState::createStream(
       ThreadState::getDefaultDevice(),
@@ -154,7 +158,27 @@ Error_t StreamCreate(Stream_t *Stream) {
   return convertAndSetLastError(Result);
 }
 
+Error_t StreamCreateWithFlags(Stream_t *Stream, unsigned int Flags) {
+  if (!Stream)
+    return setLastError(ErrorInvalidValue);
+
+  if (Flags == StreamDefault)
+    return StreamCreate(Stream);
+  if (Flags != StreamNonBlocking)
+    return setLastError(ErrorInvalidValue);
+
+  llvm::offload::StreamTy *StreamObj = nullptr;
+  ol_result_t Result = RuntimeState::createStream(
+      ThreadState::getDefaultDevice(),
+      llvm::offload::QueueKind::ExplicitNonBlocking, &StreamObj);
+  if (Result == OL_SUCCESS)
+    *Stream = makeLanguageStream(StreamObj);
+  return convertAndSetLastError(Result);
+}
+
 Error_t StreamDestroy(Stream_t Stream) {
+  if (!Stream)
+    return setLastError(ErrorInvalidValue);
   ol_result_t Result = RuntimeState::destroyStream(getInternalStream(Stream));
   return convertAndSetLastError(Result);
 }
diff --git a/offload/test/offloading/CUDA/stream_api.cu b/offload/test/offloading/CUDA/stream_api.cu
index 885ef3b300069..c5b1caaa5315c 100644
--- a/offload/test/offloading/CUDA/stream_api.cu
+++ b/offload/test/offloading/CUDA/stream_api.cu
@@ -22,12 +22,39 @@ static void print_error(const char *Label, cudaError_t Error) {
 __global__ void setValue(int *Out, int Value) { *Out = Value; }
 
 int main(int argc, char **argv) {
+  print_error("null stream create", cudaStreamCreate(nullptr));
+  // CHECK: null stream create value: 1
+  // CHECK: null stream create name: cudaErrorInvalidValue
+  print_error("null flags stream create",
+              cudaStreamCreateWithFlags(nullptr, cudaStreamDefault));
+  // CHECK: null flags stream create value: 1
+  // CHECK: null flags stream create name: cudaErrorInvalidValue
+
+  cudaStream_t InvalidFlagsStream = nullptr;
+  print_error("invalid stream flags",
+              cudaStreamCreateWithFlags(&InvalidFlagsStream, ~0u));
+  // CHECK: invalid stream flags value: 1
+  // CHECK: invalid stream flags name: cudaErrorInvalidValue
+  printf("invalid flags stream: %d\n", InvalidFlagsStream == nullptr);
+  // CHECK: invalid flags stream: 1
+
   cudaStream_t Stream = nullptr;
   if (cudaStreamCreate(&Stream) != cudaSuccess)
     return 1;
+  cudaStream_t BlockingStream = nullptr;
+  if (cudaStreamCreateWithFlags(&BlockingStream, cudaStreamDefault) !=
+      cudaSuccess)
+    return 1;
+  cudaStream_t NonBlockingStream = nullptr;
+  if (cudaStreamCreateWithFlags(&NonBlockingStream, cudaStreamNonBlocking) !=
+      cudaSuccess)
+    return 1;
 
   printf("stream created: %d\n", Stream != nullptr);
   // CHECK: stream created: 1
+  printf("stream flags created: %d %d\n", BlockingStream != nullptr,
+         NonBlockingStream != nullptr);
+  // CHECK: stream flags created: 1 1
 
   int *StreamPtr = nullptr;
   int *DefaultPtr = nullptr;
diff --git a/offload/test/offloading/HIP/stream_api.hip b/offload/test/offloading/HIP/stream_api.hip
index 6d16c2546245f..fbfca230ee697 100644
--- a/offload/test/offloading/HIP/stream_api.hip
+++ b/offload/test/offloading/HIP/stream_api.hip
@@ -22,12 +22,38 @@ static void print_error(const char *Label, hipError_t Error) {
 __global__ void setValue(int *Out, int Value) { *Out = Value; }
 
 int main(int argc, char **argv) {
+  print_error("null stream create", hipStreamCreate(nullptr));
+  // CHECK: null stream create value: 1
+  // CHECK: null stream create name: hipErrorInvalidValue
+  print_error("null flags stream create",
+              hipStreamCreateWithFlags(nullptr, hipStreamDefault));
+  // CHECK: null flags stream create value: 1
+  // CHECK: null flags stream create name: hipErrorInvalidValue
+
+  hipStream_t InvalidFlagsStream = nullptr;
+  print_error("invalid stream flags",
+              hipStreamCreateWithFlags(&InvalidFlagsStream, ~0u));
+  // CHECK: invalid stream flags value: 1
+  // CHECK: invalid stream flags name: hipErrorInvalidValue
+  printf("invalid flags stream: %d\n", InvalidFlagsStream == nullptr);
+  // CHECK: invalid flags stream: 1
+
   hipStream_t Stream = nullptr;
   if (hipStreamCreate(&Stream) != hipSuccess)
     return 1;
+  hipStream_t BlockingStream = nullptr;
+  if (hipStreamCreateWithFlags(&BlockingStream, hipStreamDefault) != hipSuccess)
+    return 1;
+  hipStream_t NonBlockingStream = nullptr;
+  if (hipStreamCreateWithFlags(&NonBlockingStream, hipStreamNonBlocking) !=
+      hipSuccess)
+    return 1;
 
   printf("stream created: %d\n", Stream != nullptr);
   // CHECK: stream created: 1
+  printf("stream flags created: %d %d\n", BlockingStream != nullptr,
+         NonBlockingStream != nullptr);
+  // CHECK: stream flags created: 1 1
 
   int *StreamPtr = nullptr;
   int *DefaultPtr = nullptr;



More information about the llvm-branch-commits mailing list