[Mlir-commits] [mlir] [mlir][acc] Don't strip a live op's implicit `acc.terminator` during host fallback (PR #205731)

llvmlistbot at llvm.org llvmlistbot at llvm.org
Wed Jun 24 23:25:51 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-mlir

Author: khaki3

<details>
<summary>Changes</summary>

Example:
```fortran
!$acc parallel if(cond)
!$acc atomic capture
  a = a + 1
  b = a
!$acc end atomic
!$acc end parallel
```

In this code, the `if` clause triggers host-fallback specialization. A blanket `acc.terminator` erase pattern can strip the implicit terminator of the still-present `acc.atomic.capture`, so the later `getTerminator()` trips `mightHaveTerminator()`.

Fix: only erase `acc.terminator` once it's no longer owned by a live ACC region op (owning ops erase their own when unwrapped). Adds a lit test.

---
Full diff: https://github.com/llvm/llvm-project/pull/205731.diff


2 Files Affected:

- (modified) mlir/lib/Dialect/OpenACC/Transforms/ACCSpecializeForHost.cpp (+24-2) 
- (modified) mlir/test/Dialect/OpenACC/acc-if-clause-lowering.mlir (+33-1) 


``````````diff
diff --git a/mlir/lib/Dialect/OpenACC/Transforms/ACCSpecializeForHost.cpp b/mlir/lib/Dialect/OpenACC/Transforms/ACCSpecializeForHost.cpp
index 633538069c268..2dfd082388873 100644
--- a/mlir/lib/Dialect/OpenACC/Transforms/ACCSpecializeForHost.cpp
+++ b/mlir/lib/Dialect/OpenACC/Transforms/ACCSpecializeForHost.cpp
@@ -264,6 +264,24 @@ class ACCOrphanAtomicCaptureOpConversion
   }
 };
 
+// Erase a stray acc.terminator only once it is no longer owned by a live ACC
+// region op; owning ops erase their own terminator when unwrapped.
+class ACCOrphanTerminatorEraseConversion
+    : public OpRewritePattern<acc::TerminatorOp> {
+  using OpRewritePattern<acc::TerminatorOp>::OpRewritePattern;
+
+public:
+  LogicalResult matchAndRewrite(acc::TerminatorOp op,
+                                PatternRewriter &rewriter) const override {
+    if (Operation *parent = op->getParentOp())
+      if (parent->getDialect() ==
+          op->getContext()->getLoadedDialect<acc::OpenACCDialect>())
+        return failure();
+    rewriter.eraseOp(op);
+    return success();
+  }
+};
+
 // Convert orphan acc.loop to scf.for or scf.execute_region.
 // Only matches if NOT inside an ACC compute construct.
 class ACCOrphanLoopOpConversion : public OpRewritePattern<acc::LoopOp> {
@@ -461,8 +479,12 @@ void mlir::acc::populateACCHostFallbackPatterns(RewritePatternSet &patterns,
   // Runtime operations - erase them
   patterns.insert<
       ACCOpEraseConversion<acc::InitOp>, ACCOpEraseConversion<acc::ShutdownOp>,
-      ACCOpEraseConversion<acc::SetOp>, ACCOpEraseConversion<acc::WaitOp>,
-      ACCOpEraseConversion<acc::TerminatorOp>>(context);
+      ACCOpEraseConversion<acc::SetOp>, ACCOpEraseConversion<acc::WaitOp>>(
+      context);
+
+  // acc.terminator - erase only stray terminators no longer owned by an ACC
+  // region op; region-bearing ops erase their own terminator when unwrapped.
+  patterns.insert<ACCOrphanTerminatorEraseConversion>(context);
 
   // Compute constructs - unwrap their regions
   patterns.insert<ACCRegionUnwrapConversion<acc::ParallelOp>,
diff --git a/mlir/test/Dialect/OpenACC/acc-if-clause-lowering.mlir b/mlir/test/Dialect/OpenACC/acc-if-clause-lowering.mlir
index 4c88df432b6c7..af4cc72a42645 100644
--- a/mlir/test/Dialect/OpenACC/acc-if-clause-lowering.mlir
+++ b/mlir/test/Dialect/OpenACC/acc-if-clause-lowering.mlir
@@ -373,4 +373,36 @@ func.func @test_acc_private(%arg0: memref<i32>, %cond: i1) {
     acc.yield
   }
   return
-}
\ No newline at end of file
+}
+
+// -----
+
+// Test that an acc.parallel with an if clause whose body holds an
+// acc.atomic.capture lowers cleanly: the device path keeps the capture and the
+// host fallback inlines it.
+// CHECK-LABEL: func.func @test_parallel_if_atomic_capture
+func.func @test_parallel_if_atomic_capture(%x: memref<i32>, %v: memref<i32>, %cond: i1) {
+  %c1_i32 = arith.constant 1 : i32
+  // CHECK-NOT: acc.parallel if
+  // CHECK: scf.if %{{.*}} {
+  // CHECK:   acc.parallel {
+  // CHECK:     acc.atomic.capture {
+  // CHECK:   } else {
+  // CHECK-NOT: acc.atomic.capture
+  // CHECK:     memref.load
+  // CHECK:     arith.addi
+  // CHECK:     memref.store
+  // CHECK:   }
+  acc.parallel if(%cond) {
+    acc.atomic.capture {
+      acc.atomic.update %x : memref<i32> {
+      ^bb0(%arg: i32):
+        %r = arith.addi %arg, %c1_i32 : i32
+        acc.yield %r : i32
+      }
+      acc.atomic.read %v = %x : memref<i32>, memref<i32>, i32
+    }
+    acc.yield
+  }
+  return
+}

``````````

</details>


https://github.com/llvm/llvm-project/pull/205731


More information about the Mlir-commits mailing list