[flang-commits] [flang] b84b33d - [flang][cuda] Add option to emit different function name for alloc/free of descriptors (#216841)

via flang-commits flang-commits at lists.llvm.org
Mon Aug 17 20:00:01 PDT 2026


Author: Valentin Clement (バレンタイン クレメン)
Date: 2026-08-17T19:59:56-07:00
New Revision: b84b33d533b1ca706c61a7fc13cda6eeaa628f55

URL: https://github.com/llvm/llvm-project/commit/b84b33d533b1ca706c61a7fc13cda6eeaa628f55
DIFF: https://github.com/llvm/llvm-project/commit/b84b33d533b1ca706c61a7fc13cda6eeaa628f55.diff

LOG: [flang][cuda] Add option to emit different function name for alloc/free of descriptors (#216841)

This allow to call specialized functions instead of the upstream ones.

Added: 
    

Modified: 
    flang/include/flang/Optimizer/Transforms/CUDA/CUFAllocationConversion.h
    flang/include/flang/Optimizer/Transforms/Passes.td
    flang/lib/Optimizer/Transforms/CUDA/CUFAllocationConversion.cpp
    flang/test/Fir/CUDA/cuda-allocate.fir

Removed: 
    


################################################################################
diff  --git a/flang/include/flang/Optimizer/Transforms/CUDA/CUFAllocationConversion.h b/flang/include/flang/Optimizer/Transforms/CUDA/CUFAllocationConversion.h
index bf34cf9117162..360cbbf3d752a 100644
--- a/flang/include/flang/Optimizer/Transforms/CUDA/CUFAllocationConversion.h
+++ b/flang/include/flang/Optimizer/Transforms/CUDA/CUFAllocationConversion.h
@@ -10,6 +10,7 @@
 #define FORTRAN_OPTIMIZER_TRANSFORMS_CUDA_CUFALLOCATIONCONVERSION_H_
 
 #include "mlir/Pass/Pass.h"
+#include "llvm/ADT/StringRef.h"
 
 namespace fir {
 class LLVMTypeConverter;
@@ -23,9 +24,14 @@ class SymbolTable;
 namespace cuf {
 
 /// Patterns that convert CUF operations to runtime calls.
+/// \p descriptorAllocFunction / \p descriptorFreeFunction, when non-empty,
+/// override the runtime functions used for descriptor allocations / frees
+/// (same signatures as CUFAllocDescriptor / CUFFreeDescriptor).
 void populateCUFAllocationConversionPatterns(
     const fir::LLVMTypeConverter &converter, mlir::DataLayout &dl,
-    const mlir::SymbolTable &symtab, mlir::RewritePatternSet &patterns);
+    const mlir::SymbolTable &symtab, mlir::RewritePatternSet &patterns,
+    llvm::StringRef descriptorAllocFunction = {},
+    llvm::StringRef descriptorFreeFunction = {});
 
 } // namespace cuf
 

diff  --git a/flang/include/flang/Optimizer/Transforms/Passes.td b/flang/include/flang/Optimizer/Transforms/Passes.td
index e7bb8ae9bb9bf..98090fefeeedc 100644
--- a/flang/include/flang/Optimizer/Transforms/Passes.td
+++ b/flang/include/flang/Optimizer/Transforms/Passes.td
@@ -529,6 +529,18 @@ def CUFAllocDelay : Pass<"cuf-alloc-delay", "::mlir::func::FuncOp"> {
 def CUFAllocationConversion : Pass<"cuf-allocation-convert", "mlir::ModuleOp"> {
   let summary = "Convert allocation related CUF operations to runtime calls";
   let dependentDialects = ["fir::FIROpsDialect"];
+  let options = [
+    Option<"descriptorAllocFunction", "descriptor-alloc-function",
+           "std::string", /*default=*/"",
+           "Name of the function to call to allocate CUDA Fortran descriptors. "
+           "Must have the same signature as CUFAllocDescriptor. "
+           "Defaults to CUFAllocDescriptor.">,
+    Option<"descriptorFreeFunction", "descriptor-free-function",
+           "std::string", /*default=*/"",
+           "Name of the function to call to free CUDA Fortran descriptors. "
+           "Must have the same signature as CUFFreeDescriptor. "
+           "Defaults to CUFFreeDescriptor.">
+  ];
 }
 
 def CUFOpConversion : Pass<"cuf-convert", "mlir::ModuleOp"> {

diff  --git a/flang/lib/Optimizer/Transforms/CUDA/CUFAllocationConversion.cpp b/flang/lib/Optimizer/Transforms/CUDA/CUFAllocationConversion.cpp
index 0032266468a6c..b75d289faca7a 100644
--- a/flang/lib/Optimizer/Transforms/CUDA/CUFAllocationConversion.cpp
+++ b/flang/lib/Optimizer/Transforms/CUDA/CUFAllocationConversion.cpp
@@ -16,6 +16,7 @@
 #include "flang/Optimizer/Dialect/FIRDialect.h"
 #include "flang/Optimizer/Dialect/FIROps.h"
 #include "flang/Optimizer/Support/DataLayout.h"
+#include "flang/Optimizer/Transforms/Passes.h"
 #include "flang/Runtime/CUDA/allocatable.h"
 #include "flang/Runtime/CUDA/common.h"
 #include "flang/Runtime/CUDA/descriptor.h"
@@ -145,12 +146,36 @@ static mlir::LogicalResult convertOpToCall(OpTy op,
   return mlir::success();
 }
 
+static mlir::func::FuncOp getDescriptorAllocFunc(mlir::Location loc,
+                                                 fir::FirOpBuilder &builder,
+                                                 llvm::StringRef customName) {
+  using RuntimeEntry = mkRTKey(CUFAllocDescriptor);
+  llvm::StringRef name = customName.empty() ? RuntimeEntry::name : customName;
+  if (auto func = builder.getNamedFunction(name))
+    return func;
+  auto funTy = RuntimeEntry::getTypeModel()(builder.getContext());
+  return builder.createRuntimeFunction(loc, name, funTy);
+}
+
+static mlir::func::FuncOp getDescriptorFreeFunc(mlir::Location loc,
+                                                fir::FirOpBuilder &builder,
+                                                llvm::StringRef customName) {
+  using RuntimeEntry = mkRTKey(CUFFreeDescriptor);
+  llvm::StringRef name = customName.empty() ? RuntimeEntry::name : customName;
+  if (auto func = builder.getNamedFunction(name))
+    return func;
+  auto funTy = RuntimeEntry::getTypeModel()(builder.getContext());
+  return builder.createRuntimeFunction(loc, name, funTy);
+}
+
 struct CUFAllocOpConversion : public mlir::OpRewritePattern<cuf::AllocOp> {
   using OpRewritePattern::OpRewritePattern;
 
   CUFAllocOpConversion(mlir::MLIRContext *context, mlir::DataLayout *dl,
-                       const fir::LLVMTypeConverter *typeConverter)
-      : OpRewritePattern(context), dl{dl}, typeConverter{typeConverter} {}
+                       const fir::LLVMTypeConverter *typeConverter,
+                       llvm::StringRef descriptorAllocFunction)
+      : OpRewritePattern(context), dl{dl}, typeConverter{typeConverter},
+        descriptorAllocFunction{descriptorAllocFunction.str()} {}
 
   mlir::LogicalResult
   matchAndRewrite(cuf::AllocOp op,
@@ -251,7 +276,7 @@ struct CUFAllocOpConversion : public mlir::OpRewritePattern<cuf::AllocOp> {
     // Convert descriptor allocations to function call.
     auto boxTy = mlir::dyn_cast_or_null<fir::BaseBoxType>(op.getInType());
     mlir::func::FuncOp func =
-        fir::runtime::getRuntimeFunc<mkRTKey(CUFAllocDescriptor)>(loc, builder);
+        getDescriptorAllocFunc(loc, builder, descriptorAllocFunction);
     auto fTy = func.getFunctionType();
     mlir::Value sourceLine =
         fir::factory::locationToLineNo(builder, loc, fTy.getInput(2));
@@ -274,11 +299,17 @@ struct CUFAllocOpConversion : public mlir::OpRewritePattern<cuf::AllocOp> {
 private:
   mlir::DataLayout *dl;
   const fir::LLVMTypeConverter *typeConverter;
+  const std::string descriptorAllocFunction;
 };
 
 struct CUFFreeOpConversion : public mlir::OpRewritePattern<cuf::FreeOp> {
   using OpRewritePattern::OpRewritePattern;
 
+  CUFFreeOpConversion(mlir::MLIRContext *context,
+                      llvm::StringRef descriptorFreeFunction)
+      : OpRewritePattern(context),
+        descriptorFreeFunction{descriptorFreeFunction.str()} {}
+
   mlir::LogicalResult
   matchAndRewrite(cuf::FreeOp op,
                   mlir::PatternRewriter &rewriter) const override {
@@ -313,7 +344,7 @@ struct CUFFreeOpConversion : public mlir::OpRewritePattern<cuf::FreeOp> {
 
     // Convert cuf.free on descriptors.
     mlir::func::FuncOp func =
-        fir::runtime::getRuntimeFunc<mkRTKey(CUFFreeDescriptor)>(loc, builder);
+        getDescriptorFreeFunc(loc, builder, descriptorFreeFunction);
     auto fTy = func.getFunctionType();
     mlir::Value sourceLine =
         fir::factory::locationToLineNo(builder, loc, fTy.getInput(2));
@@ -324,6 +355,9 @@ struct CUFFreeOpConversion : public mlir::OpRewritePattern<cuf::FreeOp> {
     rewriter.eraseOp(op);
     return mlir::success();
   }
+
+private:
+  const std::string descriptorFreeFunction;
 };
 
 struct CUFAllocateOpConversion
@@ -421,6 +455,8 @@ struct CUFDeallocateOpConversion
 class CUFAllocationConversion
     : public fir::impl::CUFAllocationConversionBase<CUFAllocationConversion> {
 public:
+  using CUFAllocationConversionBase::CUFAllocationConversionBase;
+
   void runOnOperation() override {
     auto *ctx = &getContext();
     mlir::RewritePatternSet patterns(ctx);
@@ -439,8 +475,9 @@ class CUFAllocationConversion
     target.addLegalDialect<fir::FIROpsDialect, mlir::arith::ArithDialect,
                            mlir::gpu::GPUDialect>();
     target.addLegalOp<cuf::StreamCastOp>();
-    cuf::populateCUFAllocationConversionPatterns(typeConverter, *dl, symtab,
-                                                 patterns);
+    cuf::populateCUFAllocationConversionPatterns(
+        typeConverter, *dl, symtab, patterns, descriptorAllocFunction,
+        descriptorFreeFunction);
     if (mlir::failed(mlir::applyPartialConversion(getOperation(), target,
                                                   std::move(patterns)))) {
       mlir::emitError(mlir::UnknownLoc::get(ctx),
@@ -454,8 +491,13 @@ class CUFAllocationConversion
 
 void cuf::populateCUFAllocationConversionPatterns(
     const fir::LLVMTypeConverter &converter, mlir::DataLayout &dl,
-    const mlir::SymbolTable &symtab, mlir::RewritePatternSet &patterns) {
-  patterns.insert<CUFAllocOpConversion>(patterns.getContext(), &dl, &converter);
-  patterns.insert<CUFFreeOpConversion, CUFAllocateOpConversion,
-                  CUFDeallocateOpConversion>(patterns.getContext());
+    const mlir::SymbolTable &symtab, mlir::RewritePatternSet &patterns,
+    llvm::StringRef descriptorAllocFunction,
+    llvm::StringRef descriptorFreeFunction) {
+  patterns.insert<CUFAllocOpConversion>(patterns.getContext(), &dl, &converter,
+                                        descriptorAllocFunction);
+  patterns.insert<CUFFreeOpConversion>(patterns.getContext(),
+                                       descriptorFreeFunction);
+  patterns.insert<CUFAllocateOpConversion, CUFDeallocateOpConversion>(
+      patterns.getContext());
 }

diff  --git a/flang/test/Fir/CUDA/cuda-allocate.fir b/flang/test/Fir/CUDA/cuda-allocate.fir
index c70646e312a55..e117fbe15413c 100644
--- a/flang/test/Fir/CUDA/cuda-allocate.fir
+++ b/flang/test/Fir/CUDA/cuda-allocate.fir
@@ -1,4 +1,7 @@
 // RUN: fir-opt --cuf-convert --cuf-allocation-convert %s | FileCheck %s
+// RUN: fir-opt --cuf-convert \
+// RUN:   --cuf-allocation-convert="descriptor-alloc-function=custom_alloc_desc descriptor-free-function=custom_free_desc" \
+// RUN:   %s | FileCheck %s --check-prefix=CUSTOM
 
 module attributes {dlti.dl_spec = #dlti.dl_spec<#dlti.dl_entry<f80, dense<128> : vector<2xi64>>, #dlti.dl_entry<i128, dense<128> : vector<2xi64>>, #dlti.dl_entry<i64, dense<64> : vector<2xi64>>, #dlti.dl_entry<!llvm.ptr<272>, dense<64> : vector<4xi64>>, #dlti.dl_entry<!llvm.ptr<271>, dense<32> : vector<4xi64>>, #dlti.dl_entry<!llvm.ptr<270>, dense<32> : vector<4xi64>>, #dlti.dl_entry<f128, dense<128> : vector<2xi64>>, #dlti.dl_entry<f64, dense<64> : vector<2xi64>>, #dlti.dl_entry<f16, dense<16> : vector<2xi64>>, #dlti.dl_entry<i32, dense<32> : vector<2xi64>>, #dlti.dl_entry<i16, dense<16> : vector<2xi64>>, #dlti.dl_entry<i8, dense<8> : vector<2xi64>>, #dlti.dl_entry<i1, dense<8> : vector<2xi64>>, #dlti.dl_entry<!llvm.ptr, dense<64> : vector<4xi64>>, #dlti.dl_entry<"dlti.endianness", "little">, #dlti.dl_entry<"dlti.stack_alignment", 128 : i64>>} {
 
@@ -26,6 +29,10 @@ func.func @_QPsub1() {
 // CHECK: %[[BOX_NONE:.*]] = fir.convert %[[DECL_DESC]]#1 : (!fir.ref<!fir.box<!fir.heap<!fir.array<?xf32>>>>) -> !fir.ref<!fir.box<none>>
 // CHECK: fir.call @_FortranACUFFreeDescriptor(%[[BOX_NONE]], %{{.*}}, %{{.*}}) {cuf.data_attr = #cuf.cuda<device>} : (!fir.ref<!fir.box<none>>, !fir.ref<i8>, i32) -> ()
 
+// CUSTOM-LABEL: func.func @_QPsub1()
+// CUSTOM: fir.call @custom_alloc_desc(%{{.*}}, %{{.*}}, %{{.*}}) {cuf.data_attr = #cuf.cuda<device>} : (i64, !fir.ref<i8>, i32) -> !fir.ref<!fir.box<none>>
+// CUSTOM: fir.call @custom_free_desc(%{{.*}}, %{{.*}}, %{{.*}}) {cuf.data_attr = #cuf.cuda<device>} : (!fir.ref<!fir.box<none>>, !fir.ref<i8>, i32) -> ()
+
 fir.global @_QMmod1Ea {data_attr = #cuf.cuda<device>} : !fir.box<!fir.heap<!fir.array<?xf32>>> {
     %0 = fir.zero_bits !fir.heap<!fir.array<?xf32>>
     %c0 = arith.constant 0 : index


        


More information about the flang-commits mailing list