diff --git a/enzyme/Enzyme/MLIR/Dialect/EnzymeEnums.td b/enzyme/Enzyme/MLIR/Dialect/EnzymeEnums.td index 87fa6154f15..105e7743b40 100644 --- a/enzyme/Enzyme/MLIR/Dialect/EnzymeEnums.td +++ b/enzyme/Enzyme/MLIR/Dialect/EnzymeEnums.td @@ -37,4 +37,18 @@ def ActivityAttr : EnumAttr{ }]; } +def AtomicOrdering : I32EnumAttr<"Ordering", + "Atomic ordering for LLVM's memory model", + [ + I32EnumAttrCase<"not_atomic", 0>, + I32EnumAttrCase<"unordered", 1>, + I32EnumAttrCase<"monotonic", 2>, + I32EnumAttrCase<"acquire", 4>, + I32EnumAttrCase<"release", 5>, + I32EnumAttrCase<"acq_rel", 6>, + I32EnumAttrCase<"seq_cst", 7>, + ]> { + let cppNamespace = "::mlir::enzyme"; +} + #endif // ENZYME_ENUMS diff --git a/enzyme/Enzyme/MLIR/Dialect/EnzymeOps.td b/enzyme/Enzyme/MLIR/Dialect/EnzymeOps.td index f9dc2fab17c..0808bfc9ebf 100644 --- a/enzyme/Enzyme/MLIR/Dialect/EnzymeOps.td +++ b/enzyme/Enzyme/MLIR/Dialect/EnzymeOps.td @@ -471,6 +471,38 @@ def DumpOp : Enzyme_Op<"dump"> { }]; } +def AtomicRMWOp : Enzyme_Op<"atomic_rmw", [ + AllTypesMatch<["value", "result"]>, + TypesMatchWith<"value type matches element type of memref", + "memref", "value", + "::llvm::cast($_self).getElementType()">, + ]> { + let summary = "atomic rmw operation with ordering"; + let description = [{ + }]; + + let results = (outs AnyType:$result); + + let arguments = (ins + AtomicRMWKindAttr:$kind, + AtomicOrdering:$ordering, + AnyType:$value, + Arg:$memref, + Variadic:$indices, + OptionalAttr>:$alignment); + + let assemblyFormat = [{ + $kind $value `,` $memref `[` $indices `]` $ordering attr-dict `:` `(` type($value) `,` + type($memref) `)` `->` type($result) + }]; + + let extraClassDeclaration = [{ + MemRefType getMemRefType() { + return ::llvm::cast(getMemref().getType()); + } + }]; +} + def AffineAtomicRMWOp : Enzyme_Op<"affine_atomic_rmw"> { let summary = "affine atomic rmw operation"; let description = [{ @@ -483,7 +515,8 @@ def AffineAtomicRMWOp : Enzyme_Op<"affine_atomic_rmw"> { AnyType:$value, Arg:$memref, Variadic:$indices, - AffineMapAttr:$map); + AffineMapAttr:$map, + OptionalAttr>:$alignment); let assemblyFormat = [{ $kind $value `,` $memref `,` `(` $map `)` `[` $indices `]` attr-dict `:` `(` type($value) `,` diff --git a/enzyme/Enzyme/MLIR/Implementations/AffineAutoDiffOpInterfaceImpl.cpp b/enzyme/Enzyme/MLIR/Implementations/AffineAutoDiffOpInterfaceImpl.cpp index 722720e1803..e4c02a7e364 100644 --- a/enzyme/Enzyme/MLIR/Implementations/AffineAutoDiffOpInterfaceImpl.cpp +++ b/enzyme/Enzyme/MLIR/Implementations/AffineAutoDiffOpInterfaceImpl.cpp @@ -474,20 +474,22 @@ struct AffineLoadOpInterfaceReverse } } else { bool hasIndex = loadOp.getAffineMap().getNumDims() > 0; + auto alignAttr = loadOp->getAttrOfType("alignment"); // if index had to be cached, the pop is not necessarily a valid index if (hasIndex) { SmallVector indices; computeAffineIndices(builder, loadOp.getLoc(), loadOp.getAffineMap(), retrievedArguments, indices); - memref::AtomicRMWOp::create(builder, loadOp.getLoc(), - arith::AtomicRMWKind::addf, gradient, - memrefGradient, indices); + enzyme::AtomicRMWOp::create( + builder, loadOp.getLoc(), gradient.getType(), + arith::AtomicRMWKind::addf, Ordering::monotonic, gradient, + memrefGradient, indices, alignAttr); } else { enzyme::AffineAtomicRMWOp::create( builder, loadOp.getLoc(), gradient.getType(), arith::AtomicRMWKind::addf, gradient, memrefGradient, - retrievedArguments, loadOp.getAffineMap()); + retrievedArguments, loadOp.getAffineMap(), alignAttr); } } } diff --git a/enzyme/Enzyme/MLIR/Implementations/MemRefAutoDiffOpInterfaceImpl.cpp b/enzyme/Enzyme/MLIR/Implementations/MemRefAutoDiffOpInterfaceImpl.cpp index 1a47acf6513..09611c07c7f 100644 --- a/enzyme/Enzyme/MLIR/Implementations/MemRefAutoDiffOpInterfaceImpl.cpp +++ b/enzyme/Enzyme/MLIR/Implementations/MemRefAutoDiffOpInterfaceImpl.cpp @@ -62,9 +62,10 @@ struct LoadOpInterfaceReverse memrefGradient, ArrayRef(retrievedArguments)); } else { - memref::AtomicRMWOp::create( - builder, loadOp.getLoc(), arith::AtomicRMWKind::addf, gradient, - memrefGradient, ArrayRef(retrievedArguments)); + enzyme::AtomicRMWOp::create( + builder, loadOp.getLoc(), gradient.getType(), + arith::AtomicRMWKind::addf, Ordering::monotonic, gradient, + memrefGradient, retrievedArguments, loadOp.getAlignmentAttr()); } } } diff --git a/enzyme/test/MLIR/ReverseMode/affine_parallel.mlir b/enzyme/test/MLIR/ReverseMode/affine_parallel.mlir index 3b23cbe5466..aa050185ea7 100644 --- a/enzyme/test/MLIR/ReverseMode/affine_parallel.mlir +++ b/enzyme/test/MLIR/ReverseMode/affine_parallel.mlir @@ -35,7 +35,7 @@ func.func @dfoo(%x: memref, %dx: memref, %y: memref, %dy: m // CHECK-NEXT: %4 = arith.addf %3, %cst : f32 // CHECK-NEXT: %5 = arith.mulf %2, %0 : f32 // CHECK-NEXT: %6 = arith.addf %4, %5 : f32 -// CHECK-NEXT: %7 = memref.atomic_rmw addf %6, %arg1[%arg4] : (f32, memref) -> f32 +// CHECK-NEXT: %7 = enzyme.atomic_rmw addf %6, %arg1[%arg4] monotonic : (f32, memref) -> f32 // CHECK-NEXT: } // CHECK-NEXT: memref.dealloc %alloc : memref<4xf32> // CHECK-NEXT: return @@ -79,7 +79,7 @@ func.func @dnonconst(%x: memref, %dx: memref, %y: memref, % // CHECK-NEXT: %4 = arith.addf %3, %cst : f32 // CHECK-NEXT: %5 = arith.mulf %2, %0 : f32 // CHECK-NEXT: %6 = arith.addf %4, %5 : f32 -// CHECK-NEXT: %7 = memref.atomic_rmw addf %6, %arg1[%arg5] : (f32, memref) -> f32 +// CHECK-NEXT: %7 = enzyme.atomic_rmw addf %6, %arg1[%arg5] monotonic : (f32, memref) -> f32 // CHECK-NEXT: } // CHECK-NEXT: memref.dealloc %alloc : memref // CHECK-NEXT: return @@ -125,7 +125,7 @@ func.func @dnon_1_step(%x: memref, %dx: memref, %y: memref, // CHECK-NEXT: %5 = arith.addf %4, %cst : f32 // CHECK-NEXT: %6 = arith.mulf %3, %1 : f32 // CHECK-NEXT: %7 = arith.addf %5, %6 : f32 -// CHECK-NEXT: %8 = memref.atomic_rmw addf %7, %arg1[%arg4] : (f32, memref) -> f32 +// CHECK-NEXT: %8 = enzyme.atomic_rmw addf %7, %arg1[%arg4] monotonic : (f32, memref) -> f32 // CHECK-NEXT: } // CHECK-NEXT: memref.dealloc %alloc : memref<2xf32> // CHECK-NEXT: return @@ -168,7 +168,7 @@ func.func @dpar2d(%x: memref<3x3xf32>, %dx: memref<3x3xf32>, %y: memref<3x3xf32> // CHECK-NEXT: %4 = arith.addf %3, %cst : f32 // CHECK-NEXT: %5 = arith.mulf %2, %0 : f32 // CHECK-NEXT: %6 = arith.addf %4, %5 : f32 -// CHECK-NEXT: %7 = memref.atomic_rmw addf %6, %arg1[%arg4, %arg5] : (f32, memref<3x3xf32>) -> f32 +// CHECK-NEXT: %7 = enzyme.atomic_rmw addf %6, %arg1[%arg4, %arg5] monotonic : (f32, memref<3x3xf32>) -> f32 // CHECK-NEXT: } // CHECK-NEXT: memref.dealloc %alloc : memref<3x3xf32> // CHECK-NEXT: return diff --git a/enzyme/test/MLIR/ReverseMode/affine_parallel_mincut.mlir b/enzyme/test/MLIR/ReverseMode/affine_parallel_mincut.mlir index beb560d0514..e8fbafa8310 100644 --- a/enzyme/test/MLIR/ReverseMode/affine_parallel_mincut.mlir +++ b/enzyme/test/MLIR/ReverseMode/affine_parallel_mincut.mlir @@ -110,9 +110,9 @@ func.func @dreproducer(%cond: i1, %src: memref, %dsrc: memref, %ds // CHECK: %[[ADDF_9:.*]] = arith.addf %[[ADDF_8]], %[[ADDF_7]] : f32 // CHECK: scf.yield %[[ADDF_9]], %[[ADDF_5]], %[[ADDF_6]] : f32, f32, f32 // CHECK: } -// CHECK: %[[ATOMIC_RMW_0:.*]] = memref.atomic_rmw addf %[[VAL_2:.*]]#1, %[[ARG2]]{{\[}}%[[ADDI_3]]] : (f32, memref) -> f32 -// CHECK: %[[ATOMIC_RMW_1:.*]] = memref.atomic_rmw addf %[[VAL_2]]#2, %[[ARG2]]{{\[}}%[[ADDI_2]]] : (f32, memref) -> f32 -// CHECK: %[[ATOMIC_RMW_2:.*]] = memref.atomic_rmw addf %[[VAL_2]]#0, %[[ARG2]]{{\[}}%[[VAL_1]]] : (f32, memref) -> f32 +// CHECK: %[[ATOMIC_RMW_0:.*]] = enzyme.atomic_rmw addf %[[VAL_2:.*]]#1, %[[ARG2]]{{\[}}%[[ADDI_3]]] monotonic : (f32, memref) -> f32 +// CHECK: %[[ATOMIC_RMW_1:.*]] = enzyme.atomic_rmw addf %[[VAL_2]]#2, %[[ARG2]]{{\[}}%[[ADDI_2]]] monotonic : (f32, memref) -> f32 +// CHECK: %[[ATOMIC_RMW_2:.*]] = enzyme.atomic_rmw addf %[[VAL_2]]#0, %[[ARG2]]{{\[}}%[[VAL_1]]] monotonic : (f32, memref) -> f32 // CHECK: } // CHECK: memref.dealloc %[[ALLOC_1]] : memref<100xf32> // CHECK: memref.dealloc %[[ALLOC_0]] : memref<100xf32> diff --git a/enzyme/test/MLIR/ReverseMode/parallel.mlir b/enzyme/test/MLIR/ReverseMode/parallel.mlir index a8e47001017..66983952498 100644 --- a/enzyme/test/MLIR/ReverseMode/parallel.mlir +++ b/enzyme/test/MLIR/ReverseMode/parallel.mlir @@ -43,7 +43,7 @@ module { // CHECK-NEXT: memref.store %cst, %arg4[%arg5] : memref // CHECK-NEXT: %3 = arith.mulf %2, %arg0 : f32 // CHECK-NEXT: %4 = arith.mulf %2, %1 : f32 -// CHECK-NEXT: %5 = memref.atomic_rmw addf %3, %arg2[%arg5] : (f32, memref) -> f32 +// CHECK-NEXT: %5 = enzyme.atomic_rmw addf %3, %arg2[%arg5] monotonic : (f32, memref) -> f32 // CHECK-NEXT: affine.yield %4 : f32 // CHECK-NEXT: } // CHECK-NEXT: memref.dealloc %alloc : memref<4xf32> @@ -69,7 +69,7 @@ module { // CHECK-NEXT: memref.store %cst, %arg4[%arg5] : memref // CHECK-NEXT: %3 = arith.mulf %2, %arg0 : f32 // CHECK-NEXT: %4 = arith.mulf %2, %1 : f32 -// CHECK-NEXT: %5 = memref.atomic_rmw addf %3, %arg2[%arg5] : (f32, memref) -> f32 +// CHECK-NEXT: %5 = enzyme.atomic_rmw addf %3, %arg2[%arg5] monotonic : (f32, memref) -> f32 // CHECK-NEXT: scf.reduce(%4 : f32) { // CHECK-NEXT: ^bb0(%arg6: f32, %arg7: f32): // CHECK-NEXT: %6 = arith.addf %arg6, %arg7 : f32 diff --git a/enzyme/test/MLIR/ReverseMode/scf_parallel.mlir b/enzyme/test/MLIR/ReverseMode/scf_parallel.mlir index 32abf43f76b..10af2464faf 100644 --- a/enzyme/test/MLIR/ReverseMode/scf_parallel.mlir +++ b/enzyme/test/MLIR/ReverseMode/scf_parallel.mlir @@ -42,7 +42,7 @@ func.func @dfoo(%x: memref, %dx: memref, %y: memref, %dy: m // CHECK-NEXT: %4 = arith.addf %3, %cst : f32 // CHECK-NEXT: %5 = arith.mulf %2, %0 : f32 // CHECK-NEXT: %6 = arith.addf %4, %5 : f32 -// CHECK-NEXT: %7 = memref.atomic_rmw addf %6, %arg1[%arg4] : (f32, memref) -> f32 +// CHECK-NEXT: %7 = enzyme.atomic_rmw addf %6, %arg1[%arg4] monotonic : (f32, memref) -> f32 // CHECK-NEXT: scf.reduce // CHECK-NEXT: } // CHECK-NEXT: memref.dealloc %alloc : memref<4xf32> diff --git a/enzyme/test/MLIR/ReverseMode/scf_parallel_if.mlir b/enzyme/test/MLIR/ReverseMode/scf_parallel_if.mlir index faca0abf079..8390877c11b 100644 --- a/enzyme/test/MLIR/ReverseMode/scf_parallel_if.mlir +++ b/enzyme/test/MLIR/ReverseMode/scf_parallel_if.mlir @@ -49,7 +49,7 @@ module { // CHECK: %[[x2:.+]] = memref.load %arg4[%[[arg5]]] : memref // CHECK: memref.store %[[cst]], %arg4[%[[arg5]]] : memref // CHECK: %[[x3:.+]] = arith.mulf %[[x2]], %[[x1]] : f64 - // CHECK: %[[x4:.+]] = memref.atomic_rmw addf %[[x3]], %arg1[%[[arg5]]] : (f64, memref) -> f64 + // CHECK: %[[x4:.+]] = enzyme.atomic_rmw addf %[[x3]], %arg1[%[[arg5]]] monotonic : (f64, memref) -> f64 // CHECK: scf.reduce // CHECK: } // CHECK: memref.dealloc %[[alloc]] : memref