Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 14 additions & 0 deletions enzyme/Enzyme/MLIR/Dialect/EnzymeEnums.td
Original file line number Diff line number Diff line change
Expand Up @@ -37,4 +37,18 @@ def ActivityAttr : EnumAttr<Enzyme_Dialect, Activity, "activity">{
}];
}

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
31 changes: 31 additions & 0 deletions enzyme/Enzyme/MLIR/Dialect/EnzymeOps.td
Original file line number Diff line number Diff line change
Expand Up @@ -444,6 +444,37 @@ 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<MemRefType>($_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<AnyMemRef, "the reference to rmw", [MemRead, MemWrite]>:$memref,

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we also preserve the alignment here (though I suppose we could just have alignment as a generic setAttr call)

Variadic<Index>:$indices);

let assemblyFormat = [{
$kind $value `,` $memref `[` $indices `]` $ordering attr-dict `:` `(` type($value) `,`
type($memref) `)` `->` type($result)
}];

let extraClassDeclaration = [{
MemRefType getMemRefType() {
return ::llvm::cast<MemRefType>(getMemref().getType());
}
}];
}

def AffineAtomicRMWOp : Enzyme_Op<"affine_atomic_rmw"> {
let summary = "affine atomic rmw operation";
let description = [{
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -480,9 +480,10 @@ struct AffineLoadOpInterfaceReverse
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);

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

we also need to do the same for the enzyme::AffineAtomicRMWOp too

} else {
enzyme::AffineAtomicRMWOp::create(
builder, loadOp.getLoc(), gradient.getType(),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -62,9 +62,10 @@ struct LoadOpInterfaceReverse
memrefGradient,
ArrayRef<Value>(retrievedArguments));
} else {
memref::AtomicRMWOp::create(
builder, loadOp.getLoc(), arith::AtomicRMWKind::addf, gradient,
memrefGradient, ArrayRef<Value>(retrievedArguments));
enzyme::AtomicRMWOp::create(
builder, loadOp.getLoc(), gradient.getType(),
arith::AtomicRMWKind::addf, Ordering::monotonic, gradient,
memrefGradient, retrievedArguments);
}
}
}
Expand Down
8 changes: 4 additions & 4 deletions enzyme/test/MLIR/ReverseMode/affine_parallel.mlir
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,7 @@ func.func @dfoo(%x: memref<?xf32>, %dx: memref<?xf32>, %y: memref<?xf32>, %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<?xf32>) -> f32
// CHECK-NEXT: %7 = enzyme.atomic_rmw addf %6, %arg1[%arg4] monotonic : (f32, memref<?xf32>) -> f32
// CHECK-NEXT: }
// CHECK-NEXT: memref.dealloc %alloc : memref<4xf32>
// CHECK-NEXT: return
Expand Down Expand Up @@ -79,7 +79,7 @@ func.func @dnonconst(%x: memref<?xf32>, %dx: memref<?xf32>, %y: memref<?xf32>, %
// 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<?xf32>) -> f32
// CHECK-NEXT: %7 = enzyme.atomic_rmw addf %6, %arg1[%arg5] monotonic : (f32, memref<?xf32>) -> f32
// CHECK-NEXT: }
// CHECK-NEXT: memref.dealloc %alloc : memref<?xf32>
// CHECK-NEXT: return
Expand Down Expand Up @@ -125,7 +125,7 @@ func.func @dnon_1_step(%x: memref<?xf32>, %dx: memref<?xf32>, %y: memref<?xf32>,
// 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<?xf32>) -> f32
// CHECK-NEXT: %8 = enzyme.atomic_rmw addf %7, %arg1[%arg4] monotonic : (f32, memref<?xf32>) -> f32
// CHECK-NEXT: }
// CHECK-NEXT: memref.dealloc %alloc : memref<2xf32>
// CHECK-NEXT: return
Expand Down Expand Up @@ -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
Expand Down
6 changes: 3 additions & 3 deletions enzyme/test/MLIR/ReverseMode/affine_parallel_mincut.mlir
Original file line number Diff line number Diff line change
Expand Up @@ -110,9 +110,9 @@ func.func @dreproducer(%cond: i1, %src: memref<?xf32>, %dsrc: memref<?xf32>, %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<?xf32>) -> f32
// CHECK: %[[ATOMIC_RMW_1:.*]] = memref.atomic_rmw addf %[[VAL_2]]#2, %[[ARG2]]{{\[}}%[[ADDI_2]]] : (f32, memref<?xf32>) -> f32
// CHECK: %[[ATOMIC_RMW_2:.*]] = memref.atomic_rmw addf %[[VAL_2]]#0, %[[ARG2]]{{\[}}%[[VAL_1]]] : (f32, memref<?xf32>) -> f32
// CHECK: %[[ATOMIC_RMW_0:.*]] = enzyme.atomic_rmw addf %[[VAL_2:.*]]#1, %[[ARG2]]{{\[}}%[[ADDI_3]]] monotonic : (f32, memref<?xf32>) -> f32
// CHECK: %[[ATOMIC_RMW_1:.*]] = enzyme.atomic_rmw addf %[[VAL_2]]#2, %[[ARG2]]{{\[}}%[[ADDI_2]]] monotonic : (f32, memref<?xf32>) -> f32
// CHECK: %[[ATOMIC_RMW_2:.*]] = enzyme.atomic_rmw addf %[[VAL_2]]#0, %[[ARG2]]{{\[}}%[[VAL_1]]] monotonic : (f32, memref<?xf32>) -> f32
// CHECK: }
// CHECK: memref.dealloc %[[ALLOC_1]] : memref<100xf32>
// CHECK: memref.dealloc %[[ALLOC_0]] : memref<100xf32>
Expand Down
4 changes: 2 additions & 2 deletions enzyme/test/MLIR/ReverseMode/parallel.mlir
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,7 @@ module {
// CHECK-NEXT: memref.store %cst, %arg4[%arg5] : memref<?xf32>
// 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<?xf32>) -> f32
// CHECK-NEXT: %5 = enzyme.atomic_rmw addf %3, %arg2[%arg5] monotonic : (f32, memref<?xf32>) -> f32
// CHECK-NEXT: affine.yield %4 : f32
// CHECK-NEXT: }
// CHECK-NEXT: memref.dealloc %alloc : memref<4xf32>
Expand All @@ -69,7 +69,7 @@ module {
// CHECK-NEXT: memref.store %cst, %arg4[%arg5] : memref<?xf32>
// 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<?xf32>) -> f32
// CHECK-NEXT: %5 = enzyme.atomic_rmw addf %3, %arg2[%arg5] monotonic : (f32, memref<?xf32>) -> f32
// CHECK-NEXT: scf.reduce(%4 : f32) {
// CHECK-NEXT: ^bb0(%arg6: f32, %arg7: f32):
// CHECK-NEXT: %6 = arith.addf %arg6, %arg7 : f32
Expand Down
2 changes: 1 addition & 1 deletion enzyme/test/MLIR/ReverseMode/scf_parallel.mlir
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,7 @@ func.func @dfoo(%x: memref<?xf32>, %dx: memref<?xf32>, %y: memref<?xf32>, %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<?xf32>) -> f32
// CHECK-NEXT: %7 = enzyme.atomic_rmw addf %6, %arg1[%arg4] monotonic : (f32, memref<?xf32>) -> f32
// CHECK-NEXT: scf.reduce
// CHECK-NEXT: }
// CHECK-NEXT: memref.dealloc %alloc : memref<4xf32>
Expand Down
2 changes: 1 addition & 1 deletion enzyme/test/MLIR/ReverseMode/scf_parallel_if.mlir
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,7 @@ module {
// CHECK: %[[x2:.+]] = memref.load %arg4[%[[arg5]]] : memref<?xf64>
// CHECK: memref.store %[[cst]], %arg4[%[[arg5]]] : memref<?xf64>
// CHECK: %[[x3:.+]] = arith.mulf %[[x2]], %[[x1]] : f64
// CHECK: %[[x4:.+]] = memref.atomic_rmw addf %[[x3]], %arg1[%[[arg5]]] : (f64, memref<?xf64>) -> f64
// CHECK: %[[x4:.+]] = enzyme.atomic_rmw addf %[[x3]], %arg1[%[[arg5]]] monotonic : (f64, memref<?xf64>) -> f64
// CHECK: scf.reduce
// CHECK: }
// CHECK: memref.dealloc %[[alloc]] : memref<?xf64>
Expand Down
Loading