[flang-commits] [flang] [flang][cuda] Synchronize context before finalization (PR #223058)

Valentin Clement バレンタイン クレメン via flang-commits flang-commits at lists.llvm.org
Fri Sep 11 13:55:38 PDT 2026


https://github.com/clementval updated https://github.com/llvm/llvm-project/pull/223058

>From 6fe4ad6eb9cc69f432525d353ddfaf62410c8bab Mon Sep 17 00:00:00 2001
From: Valentin Clement <clementval at gmail.com>
Date: Fri, 11 Sep 2026 13:50:15 -0700
Subject: [PATCH 1/2] [flang][cuda] Synchronize context before finalization

---
 .../Optimizer/Builder/Runtime/CUDA/Support.h  | 27 +++++++++++++++
 flang/lib/Lower/Bridge.cpp                    |  6 +++-
 flang/lib/Optimizer/Builder/CMakeLists.txt    |  1 +
 .../Builder/Runtime/CUDA/Support.cpp          | 33 +++++++++++++++++++
 flang/test/Lower/CUDA/cuda-return01.cuf       |  1 +
 flang/test/Lower/CUDA/cuda-return02.cuf       |  3 ++
 6 files changed, 70 insertions(+), 1 deletion(-)
 create mode 100644 flang/include/flang/Optimizer/Builder/Runtime/CUDA/Support.h
 create mode 100644 flang/lib/Optimizer/Builder/Runtime/CUDA/Support.cpp

diff --git a/flang/include/flang/Optimizer/Builder/Runtime/CUDA/Support.h b/flang/include/flang/Optimizer/Builder/Runtime/CUDA/Support.h
new file mode 100644
index 0000000000000..c22b61b854017
--- /dev/null
+++ b/flang/include/flang/Optimizer/Builder/Runtime/CUDA/Support.h
@@ -0,0 +1,27 @@
+//===-- Support.h - CUDA support runtime functions --------------*- C++ -*-===//
+//
+// 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 FORTRAN_OPTIMIZER_BUILDER_RUNTIME_CUDA_SUPPORT_H_
+#define FORTRAN_OPTIMIZER_BUILDER_RUNTIME_CUDA_SUPPORT_H_
+
+namespace mlir {
+class Location;
+} // namespace mlir
+
+namespace fir {
+class FirOpBuilder;
+}
+
+namespace fir::runtime::cuda {
+
+/// Generate runtime call to synchronize the CUDA device.
+void getCUDADeviceSynchronize(fir::FirOpBuilder &builder, mlir::Location loc);
+
+} // namespace fir::runtime::cuda
+
+#endif // FORTRAN_OPTIMIZER_BUILDER_RUNTIME_CUDA_SUPPORT_H_
diff --git a/flang/lib/Lower/Bridge.cpp b/flang/lib/Lower/Bridge.cpp
index 5ea1850669297..b6b8b03073ce3 100644
--- a/flang/lib/Lower/Bridge.cpp
+++ b/flang/lib/Lower/Bridge.cpp
@@ -39,6 +39,7 @@
 #include "flang/Optimizer/Builder/FIRBuilder.h"
 #include "flang/Optimizer/Builder/Runtime/Assign.h"
 #include "flang/Optimizer/Builder/Runtime/CUDA/Descriptor.h"
+#include "flang/Optimizer/Builder/Runtime/CUDA/Support.h"
 #include "flang/Optimizer/Builder/Runtime/Character.h"
 #include "flang/Optimizer/Builder/Runtime/Derived.h"
 #include "flang/Optimizer/Builder/Runtime/EnvironmentDefaults.h"
@@ -2005,7 +2006,10 @@ class FirConverter : public Fortran::lower::AbstractConverter {
         mlir::Value active =
             fir::runtime::cuda::genDeviceIsActive(*builder, loc);
         builder->genIfThen(loc, active)
-            .genThen([&]() { bridge.cudaCleanupCtx().finalizeAndKeep(); })
+            .genThen([&]() { 
+              fir::runtime::cuda::getCUDADeviceSynchronize(*builder, loc);
+              bridge.cudaCleanupCtx().finalizeAndKeep(); 
+            })
             .end();
       }
       bridge.fctCtx().finalizeAndKeep();
diff --git a/flang/lib/Optimizer/Builder/CMakeLists.txt b/flang/lib/Optimizer/Builder/CMakeLists.txt
index b77e21d40476c..190b16d659338 100644
--- a/flang/lib/Optimizer/Builder/CMakeLists.txt
+++ b/flang/lib/Optimizer/Builder/CMakeLists.txt
@@ -22,6 +22,7 @@ add_flang_library(FIRBuilder
   Runtime/Character.cpp
   Runtime/Command.cpp
   Runtime/CUDA/Descriptor.cpp
+  Runtime/CUDA/Support.cpp
   Runtime/Derived.cpp
   Runtime/EnvironmentDefaults.cpp
   Runtime/Exceptions.cpp
diff --git a/flang/lib/Optimizer/Builder/Runtime/CUDA/Support.cpp b/flang/lib/Optimizer/Builder/Runtime/CUDA/Support.cpp
new file mode 100644
index 0000000000000..4fc5fa8e3c090
--- /dev/null
+++ b/flang/lib/Optimizer/Builder/Runtime/CUDA/Support.cpp
@@ -0,0 +1,33 @@
+//===-- Support.cpp -- Lowering helper for CUDA runtime functions ---------===//
+//
+// 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
+//
+//===----------------------------------------------------------------------===//
+
+#include "flang/Optimizer/Builder/Runtime/CUDA/Support.h"
+#include "flang/Optimizer/Builder/FIRBuilder.h"
+#include "flang/Optimizer/Builder/Runtime/RTBuilder.h"
+
+using namespace fir::runtime::cuda;
+
+static constexpr llvm::StringRef kCudaDeviceSynchronizeName = "_QPcudadevicesynchronize";
+
+void fir::runtime::cuda::getCUDADeviceSynchronize(fir::FirOpBuilder &builder,
+                                                  mlir::Location loc) {
+  mlir::func::FuncOp func = builder.getNamedFunction(kCudaDeviceSynchronizeName);
+  if (!func) {
+    mlir::FunctionType funcType = mlir::FunctionType::get(builder.getContext(), {}, {builder.getI32Type()});
+    func = builder.createFunction(loc, kCudaDeviceSynchronizeName, funcType);
+    func->setAttr(fir::getFortranProcedureFlagsAttrName(), fir::FortranProcedureFlagsEnumAttr::get(builder.getContext(), fir::FortranProcedureFlagsEnum::intrinsic));
+    func.setPrivate();
+  }
+  auto call = fir::CallOp::create(builder, loc, func, mlir::ValueRange{});
+  call.setProcedureAttrsAttr(fir::FortranProcedureFlagsEnumAttr::get(builder.getContext(), fir::FortranProcedureFlagsEnum::intrinsic));
+}
+
+
+//  %1315 = fir.call @_QPcudadevicesynchronize() proc_attrs<intrinsic> fastmath<contract> : () -> i32
+
+//  func.func private @_QPcudadevicesynchronize() -> i32 attributes {fir.proc_attrs = #fir.proc_attrs<intrinsic>}
diff --git a/flang/test/Lower/CUDA/cuda-return01.cuf b/flang/test/Lower/CUDA/cuda-return01.cuf
index 9ea5ce1538081..847a22cc18229 100644
--- a/flang/test/Lower/CUDA/cuda-return01.cuf
+++ b/flang/test/Lower/CUDA/cuda-return01.cuf
@@ -43,6 +43,7 @@ end
 ! CHECK: %[[PTR:.*]]:2 = hlfir.declare %{{.*}} {data_attr = #cuf.cuda<device>, fortran_attrs = #fir.var_attrs<pointer>, uniq_name = "_QFEp"}
 ! CHECK: %[[ACTIVE:.*]] = fir.call @_FortranACUFDeviceIsActive() {{.*}} : () -> i1
 ! CHECK-NEXT: fir.if %[[ACTIVE]] {
+! CHECK-NEXT: fir.call @_QPcudadevicesynchronize()
 ! CHECK-NEXT: cuf.free %[[PTR]]#0 : !fir.ref<!fir.box<!fir.ptr<!fir.array<?xi32>>>>{{.*}}
 ! CHECK: cuf.deallocate
 ! CHECK: cuf.free{{.*}}
diff --git a/flang/test/Lower/CUDA/cuda-return02.cuf b/flang/test/Lower/CUDA/cuda-return02.cuf
index b45635c219202..aa82cd5039cb2 100644
--- a/flang/test/Lower/CUDA/cuda-return02.cuf
+++ b/flang/test/Lower/CUDA/cuda-return02.cuf
@@ -19,12 +19,14 @@ end
 ! CHECK-NEXT: ^bb1:
 ! CHECK-NEXT: %[[ACTIVE1:.*]] = fir.call @_FortranACUFDeviceIsActive()
 ! CHECK-NEXT: fir.if %[[ACTIVE1]] {
+! CHECK-NEXT: fir.call @_QPcudadevicesynchronize() proc_attrs<intrinsic> fastmath<contract> : () -> i32
 ! CHECK-NEXT: cuf.free %[[DECL]]#0 : !fir.ref<!fir.array<10xi32>>{{.*}}
 ! CHECK-NEXT: }
 ! CHECK-NEXT: return
 ! CHECK-NEXT: ^bb2:
 ! CHECK-NEXT: %[[ACTIVE2:.*]] = fir.call @_FortranACUFDeviceIsActive()
 ! CHECK-NEXT: fir.if %[[ACTIVE2]] {
+! CHECK-NEXT: fir.call @_QPcudadevicesynchronize() proc_attrs<intrinsic> fastmath<contract> : () -> i32
 ! CHECK-NEXT: cuf.free %[[DECL]]#0 : !fir.ref<!fir.array<10xi32>>{{.*}}
 ! CHECK-NEXT: }
 ! CHECK-NEXT: return
@@ -52,6 +54,7 @@ end
 ! CHECK: ^bb3:
 ! CHECK-NEXT: %[[SUB_ACTIVE:.*]] = fir.call @_FortranACUFDeviceIsActive()
 ! CHECK-NEXT: fir.if %[[SUB_ACTIVE]] {
+! CHECK-NEXT: fir.call @_QPcudadevicesynchronize() proc_attrs<intrinsic> fastmath<contract> : () -> i32
 ! CHECK-NEXT: cuf.free %[[DECL]]#0 : !fir.ref<!fir.array<10xi32>>{{.*}}
 ! CHECK-NEXT: }
 ! CHECK-NEXT: return

>From c001e971b3f931d929e6eb7ac1ce98cb7bf539f1 Mon Sep 17 00:00:00 2001
From: Valentin Clement <clementval at gmail.com>
Date: Fri, 11 Sep 2026 13:55:14 -0700
Subject: [PATCH 2/2] format

---
 flang/lib/Lower/Bridge.cpp                    |  4 ++--
 .../Builder/Runtime/CUDA/Support.cpp          | 24 ++++++++++++-------
 2 files changed, 18 insertions(+), 10 deletions(-)

diff --git a/flang/lib/Lower/Bridge.cpp b/flang/lib/Lower/Bridge.cpp
index b6b8b03073ce3..30e7903966ea4 100644
--- a/flang/lib/Lower/Bridge.cpp
+++ b/flang/lib/Lower/Bridge.cpp
@@ -2006,9 +2006,9 @@ class FirConverter : public Fortran::lower::AbstractConverter {
         mlir::Value active =
             fir::runtime::cuda::genDeviceIsActive(*builder, loc);
         builder->genIfThen(loc, active)
-            .genThen([&]() { 
+            .genThen([&]() {
               fir::runtime::cuda::getCUDADeviceSynchronize(*builder, loc);
-              bridge.cudaCleanupCtx().finalizeAndKeep(); 
+              bridge.cudaCleanupCtx().finalizeAndKeep();
             })
             .end();
       }
diff --git a/flang/lib/Optimizer/Builder/Runtime/CUDA/Support.cpp b/flang/lib/Optimizer/Builder/Runtime/CUDA/Support.cpp
index 4fc5fa8e3c090..0973178576712 100644
--- a/flang/lib/Optimizer/Builder/Runtime/CUDA/Support.cpp
+++ b/flang/lib/Optimizer/Builder/Runtime/CUDA/Support.cpp
@@ -12,22 +12,30 @@
 
 using namespace fir::runtime::cuda;
 
-static constexpr llvm::StringRef kCudaDeviceSynchronizeName = "_QPcudadevicesynchronize";
+static constexpr llvm::StringRef kCudaDeviceSynchronizeName =
+    "_QPcudadevicesynchronize";
 
 void fir::runtime::cuda::getCUDADeviceSynchronize(fir::FirOpBuilder &builder,
                                                   mlir::Location loc) {
-  mlir::func::FuncOp func = builder.getNamedFunction(kCudaDeviceSynchronizeName);
+  mlir::func::FuncOp func =
+      builder.getNamedFunction(kCudaDeviceSynchronizeName);
   if (!func) {
-    mlir::FunctionType funcType = mlir::FunctionType::get(builder.getContext(), {}, {builder.getI32Type()});
+    mlir::FunctionType funcType = mlir::FunctionType::get(
+        builder.getContext(), {}, {builder.getI32Type()});
     func = builder.createFunction(loc, kCudaDeviceSynchronizeName, funcType);
-    func->setAttr(fir::getFortranProcedureFlagsAttrName(), fir::FortranProcedureFlagsEnumAttr::get(builder.getContext(), fir::FortranProcedureFlagsEnum::intrinsic));
+    func->setAttr(
+        fir::getFortranProcedureFlagsAttrName(),
+        fir::FortranProcedureFlagsEnumAttr::get(
+            builder.getContext(), fir::FortranProcedureFlagsEnum::intrinsic));
     func.setPrivate();
   }
   auto call = fir::CallOp::create(builder, loc, func, mlir::ValueRange{});
-  call.setProcedureAttrsAttr(fir::FortranProcedureFlagsEnumAttr::get(builder.getContext(), fir::FortranProcedureFlagsEnum::intrinsic));
+  call.setProcedureAttrsAttr(fir::FortranProcedureFlagsEnumAttr::get(
+      builder.getContext(), fir::FortranProcedureFlagsEnum::intrinsic));
 }
 
+//  %1315 = fir.call @_QPcudadevicesynchronize() proc_attrs<intrinsic>
+//  fastmath<contract> : () -> i32
 
-//  %1315 = fir.call @_QPcudadevicesynchronize() proc_attrs<intrinsic> fastmath<contract> : () -> i32
-
-//  func.func private @_QPcudadevicesynchronize() -> i32 attributes {fir.proc_attrs = #fir.proc_attrs<intrinsic>}
+//  func.func private @_QPcudadevicesynchronize() -> i32 attributes
+//  {fir.proc_attrs = #fir.proc_attrs<intrinsic>}



More information about the flang-commits mailing list