[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