[Mlir-commits] [mlir] d2c003c - [mlir][OpenACC] Lower constant sized loop clauses (#219043)

llvmlistbot at llvm.org llvmlistbot at llvm.org
Thu Aug 27 13:56:26 PDT 2026


Author: Delaram Talaashrafi
Date: 2026-08-27T16:56:20-04:00
New Revision: d2c003c4d2c863eb719071679ee9c41fdba19b7e

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

LOG: [mlir][OpenACC] Lower constant sized loop clauses (#219043)

Add support for sized loop clauses (`acc.loop
vector(n)/worker(n)/gang(num:n)`) in kernels constructs. Collect
constant sizes from kernels loops before conversion and add them to the
launch arguments. Treat the parallelism levels associated with sized
clauses as regular levels when assigning par_dims to loops. Non-constant
sizes remain NYI.

Added: 
    

Modified: 
    mlir/lib/Dialect/OpenACC/Transforms/ACCComputeLowering.cpp
    mlir/test/Dialect/OpenACC/acc-compute-lowering-compute.mlir

Removed: 
    


################################################################################
diff  --git a/mlir/lib/Dialect/OpenACC/Transforms/ACCComputeLowering.cpp b/mlir/lib/Dialect/OpenACC/Transforms/ACCComputeLowering.cpp
index a5fbed1d36fd7..832de220289ea 100644
--- a/mlir/lib/Dialect/OpenACC/Transforms/ACCComputeLowering.cpp
+++ b/mlir/lib/Dialect/OpenACC/Transforms/ACCComputeLowering.cpp
@@ -47,6 +47,7 @@
 
 #include "mlir/Dialect/Arith/IR/Arith.h"
 #include "mlir/Dialect/Func/IR/FuncOps.h"
+#include "mlir/Dialect/OpenACC/Analysis/OpenACCSupport.h"
 #include "mlir/Dialect/OpenACC/OpenACC.h"
 #include "mlir/Dialect/OpenACC/OpenACCParMapping.h"
 #include "mlir/Dialect/OpenACC/OpenACCUtils.h"
@@ -58,6 +59,7 @@
 #include "mlir/Interfaces/FunctionInterfaces.h"
 #include "mlir/Transforms/GreedyPatternRewriteDriver.h"
 #include "mlir/Transforms/RegionUtils.h"
+#include "llvm/ADT/DenseMap.h"
 #include "llvm/ADT/STLExtras.h"
 
 namespace mlir {
@@ -153,8 +155,59 @@ static DeviceType getParDimsDeviceType(ComputeConstructT computeOp,
   return DeviceType::None;
 }
 
+/// Constant sized gang/worker/vector clauses collected per compute construct.
+struct SizedLevel {
+  ParLevel level;
+  int64_t size;
+};
+using SizedLevelMap = DenseMap<Operation *, SmallVector<SizedLevel>>;
+
+/// Record a sized clause if `size` is a constant; NYI otherwise.
+static LogicalResult tryAddSizedLevel(SizedLevelMap &sizedLevelMap,
+                                      Operation *computeOp, LoopOp loopOp,
+                                      ParLevel level, Value size,
+                                      OpenACCSupport &accSupport) {
+  if (!size)
+    return success();
+  std::optional<int64_t> constSize = getConstantIntValue(size);
+  if (!constSize) {
+    accSupport.emitNYI(loopOp.getLoc(),
+                       "non-constant sized parallelism clause");
+    return failure();
+  }
+  sizedLevelMap[computeOp].push_back({level, *constSize});
+  return success();
+}
+
+/// Collect constant sized levels from loops in `acc.kernels` regions.
+static LogicalResult fillSizedLevelMap(Operation *op, DeviceType deviceType,
+                                       SizedLevelMap &sizedLevelMap,
+                                       OpenACCSupport &accSupport) {
+  WalkResult result = op->walk([&](LoopOp loopOp) {
+    Operation *computeOp =
+        getEnclosingComputeOp(*loopOp->getBlock()->getParent());
+    if (!computeOp || !isa<KernelsOp>(computeOp))
+      return WalkResult::advance();
+    DeviceType loopDeviceType =
+        getGangWorkerVectorDeviceType(loopOp, deviceType);
+    if (failed(tryAddSizedLevel(
+            sizedLevelMap, computeOp, loopOp, ParLevel::vector,
+            loopOp.getVectorValue(loopDeviceType), accSupport)) ||
+        failed(tryAddSizedLevel(
+            sizedLevelMap, computeOp, loopOp, ParLevel::worker,
+            loopOp.getWorkerValue(loopDeviceType), accSupport)) ||
+        failed(tryAddSizedLevel(
+            sizedLevelMap, computeOp, loopOp, ParLevel::gang_dim1,
+            loopOp.getGangValue(GangArgType::Num, loopDeviceType), accSupport)))
+      return WalkResult::interrupt();
+    return WalkResult::advance();
+  });
+  return failure(result.wasInterrupted());
+}
+
 /// Map loop parallelism clauses (gang/worker/vector) to GPU parallel
-/// dimensions using the given mapping policy.
+/// dimensions using the given mapping policy. Sized clauses (e.g. vector(n))
+/// count as the corresponding level.
 static SmallVector<GPUParallelDimAttr>
 getParallelDimensions(LoopOp loopOp, const ACCToGPUMappingPolicy &policy,
                       DeviceType deviceType) {
@@ -162,9 +215,9 @@ getParallelDimensions(LoopOp loopOp, const ACCToGPUMappingPolicy &policy,
   SmallVector<GPUParallelDimAttr> parDims;
   auto *ctx = loopOp->getContext();
 
-  if (loopOp.hasVector(deviceType))
+  if (loopOp.hasVector(deviceType) || loopOp.getVectorValue(deviceType))
     insertParDim(parDims, policy.vectorDim(ctx));
-  if (loopOp.hasWorker(deviceType))
+  if (loopOp.hasWorker(deviceType) || loopOp.getWorkerValue(deviceType))
     insertParDim(parDims, policy.workerDim(ctx));
   if (auto gangDimValue = loopOp.getGangValue(GangArgType::Dim, deviceType)) {
     if (auto gangDimDefOp =
@@ -172,7 +225,8 @@ getParallelDimensions(LoopOp loopOp, const ACCToGPUMappingPolicy &policy,
       auto gangLevel = getGangParLevel(gangDimDefOp.value());
       insertParDim(parDims, policy.gangDim(ctx, gangLevel));
     }
-  } else if (loopOp.hasGang(deviceType)) {
+  } else if (loopOp.hasGang(deviceType) ||
+             loopOp.getGangValue(GangArgType::Num, deviceType)) {
     insertParDim(parDims, policy.gangDim(ctx, ParLevel::gang_dim1));
   }
   return parDims;
@@ -184,10 +238,9 @@ getParallelDimensions(LoopOp loopOp, const ACCToGPUMappingPolicy &policy,
 /// `acc.par_width` from gang/worker/vector (device-type operands first, then
 /// default DeviceType::None).
 template <typename ComputeConstructT>
-static SmallVector<Value>
-assignKnownLaunchArgs(ComputeConstructT computeOp, DeviceType deviceType,
-                      RewriterBase &rewriter,
-                      const ACCToGPUMappingPolicy &policy) {
+static SmallVector<Value> assignKnownLaunchArgs(
+    ComputeConstructT computeOp, DeviceType deviceType, RewriterBase &rewriter,
+    const ACCToGPUMappingPolicy &policy, const SizedLevelMap &sizedLevelMap) {
   auto *ctx = rewriter.getContext();
   auto loc = computeOp->getLoc();
 
@@ -230,6 +283,24 @@ assignKnownLaunchArgs(ComputeConstructT computeOp, DeviceType deviceType,
                                           indexTy, vectorLength),
           policy.vectorDim(ctx)));
     }
+
+    // Loop-level sized clauses. Skip a dim already set on the construct.
+    // Rematerialize the constant here so it dominates the compute region.
+    auto sizedLevels = sizedLevelMap.find(computeOp.getOperation());
+    if (sizedLevels != sizedLevelMap.end()) {
+      for (const SizedLevel &sizedLevel : sizedLevels->second) {
+        GPUParallelDimAttr dim = policy.map(ctx, sizedLevel.level);
+        bool exists = llvm::any_of(values, [&](Value v) {
+          auto parWidth = v.getDefiningOp<ParWidthOp>();
+          return parWidth && parWidth.getParDim() == dim;
+        });
+        if (exists)
+          continue;
+        Value sizeVal =
+            arith::ConstantIndexOp::create(rewriter, loc, sizedLevel.size);
+        values.push_back(ParWidthOp::create(rewriter, loc, sizeVal, dim));
+      }
+    }
     return values;
   } else {
     llvm_unreachable("assignKnownLaunchArgs: expected parallel, kernels, or "
@@ -329,17 +400,17 @@ template <typename ComputeConstructT>
 class ComputeOpConversion : public OpRewritePattern<ComputeConstructT> {
 public:
   ComputeOpConversion(MLIRContext *ctx, const ACCToGPUMappingPolicy &policy,
-                      DeviceType deviceType)
+                      DeviceType deviceType, const SizedLevelMap &sizedLevelMap)
       : OpRewritePattern<ComputeConstructT>(ctx), policy(policy),
-        deviceType(deviceType) {}
+        deviceType(deviceType), sizedLevelMap(sizedLevelMap) {}
 
   LogicalResult matchAndRewrite(ComputeConstructT computeOp,
                                 PatternRewriter &rewriter) const override {
     rewriter.setInsertionPoint(computeOp);
     auto kernelEnv =
         KernelEnvironmentOp::createAndPopulate(computeOp, deviceType, rewriter);
-    auto launchArgs =
-        assignKnownLaunchArgs(computeOp, deviceType, rewriter, policy);
+    auto launchArgs = assignKnownLaunchArgs(computeOp, deviceType, rewriter,
+                                            policy, sizedLevelMap);
     Region &region = computeOp.getRegion();
     SetVector<Value> liveInValues;
     getUsedValuesDefinedAbove(region, region, liveInValues);
@@ -359,6 +430,7 @@ class ComputeOpConversion : public OpRewritePattern<ComputeConstructT> {
 private:
   const ACCToGPUMappingPolicy &policy;
   DeviceType deviceType;
+  const SizedLevelMap &sizedLevelMap;
 };
 
 //===----------------------------------------------------------------------===//
@@ -375,6 +447,11 @@ class ACCComputeLowering
     auto *context = op.getContext();
 
     DefaultACCToGPUMappingPolicy policy;
+    // Collect loop sized levels before loops are rewritten away.
+    SizedLevelMap sizedLevelMap;
+    OpenACCSupport &accSupport = getAnalysis<OpenACCSupport>();
+    if (failed(fillSizedLevelMap(op, deviceType, sizedLevelMap, accSupport)))
+      return signalPassFailure();
 
     // Part 1: Convert acc.loop to scf.parallel/scf.for while the parent
     // compute construct is still present (needed to determine conversion
@@ -389,7 +466,8 @@ class ACCComputeLowering
     RewritePatternSet computePatterns(context);
     computePatterns
         .insert<ComputeOpConversion<ParallelOp>, ComputeOpConversion<KernelsOp>,
-                ComputeOpConversion<SerialOp>>(context, policy, deviceType);
+                ComputeOpConversion<SerialOp>>(context, policy, deviceType,
+                                               sizedLevelMap);
     if (failed(applyPatternsGreedily(op, std::move(computePatterns))))
       return signalPassFailure();
   }

diff  --git a/mlir/test/Dialect/OpenACC/acc-compute-lowering-compute.mlir b/mlir/test/Dialect/OpenACC/acc-compute-lowering-compute.mlir
index b85d36914fc1e..f871932587495 100644
--- a/mlir/test/Dialect/OpenACC/acc-compute-lowering-compute.mlir
+++ b/mlir/test/Dialect/OpenACC/acc-compute-lowering-compute.mlir
@@ -270,3 +270,34 @@ func.func @parallel_num_gangs_1_2_independent(%buf: memref<4xi32>) {
   acc.copyout accPtr(%dev : memref<4xi32>) to varPtr(%buf : memref<4xi32>)
   return
 }
+
+// -----
+
+// A sized `vector(n)` clause on a loop inside acc.kernels supplies the vector
+// launch width when the construct itself has no vector_length clause.
+
+// CHECK-LABEL: func.func @kernels_loop_sized_vector
+func.func @kernels_loop_sized_vector(%buf: memref<4xi32>) {
+  %c0 = arith.constant 0 : index
+  %c1 = arith.constant 1 : index
+  %c4 = arith.constant 4 : index
+  %c32_i32 = arith.constant 32 : i32
+
+  %dev = acc.copyin varPtr(%buf : memref<4xi32>) -> memref<4xi32>
+  // CHECK: %[[VL:.*]] = arith.constant 32 : index
+  // CHECK: acc.kernel_environment
+  // CHECK: %[[PW:.*]] = acc.par_width %[[VL]] par_dim(#acc.par_dim<thread_x>)
+  // CHECK: acc.compute_region launch(%{{.*}} = %[[PW]])
+  // CHECK: scf.parallel
+  // CHECK: acc.par_dims = #acc<par_dims[thread_x]>
+  acc.kernels dataOperands(%dev : memref<4xi32>) {
+    acc.loop vector(%c32_i32 : i32) control(%i : index) = (%c0 : index) to (%c4 : index) step (%c1 : index) {
+      %vi = arith.index_cast %i : index to i32
+      memref.store %vi, %dev[%i] : memref<4xi32>
+      acc.yield
+    } independent
+    acc.terminator
+  }
+  acc.copyout accPtr(%dev : memref<4xi32>) to varPtr(%buf : memref<4xi32>)
+  return
+}


        


More information about the Mlir-commits mailing list