Skip to content
Merged
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
35 changes: 34 additions & 1 deletion enzyme/Enzyme/MLIR/Dialect/EnzymeOps.td
Original file line number Diff line number Diff line change
Expand Up @@ -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<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,
Comment thread
wsmoses marked this conversation as resolved.
Variadic<Index>:$indices,
OptionalAttr<IntValidAlignment<I64Attr>>:$alignment);

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 All @@ -483,7 +515,8 @@ def AffineAtomicRMWOp : Enzyme_Op<"affine_atomic_rmw"> {
AnyType:$value,
Arg<AnyMemRef, "the reference to rmw", [MemRead, MemWrite]>:$memref,
Variadic<Index>:$indices,
AffineMapAttr:$map);
AffineMapAttr:$map,
OptionalAttr<IntValidAlignment<I64Attr>>:$alignment);

let assemblyFormat = [{
$kind $value `,` $memref `,` `(` $map `)` `[` $indices `]` attr-dict `:` `(` type($value) `,`
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -474,20 +474,22 @@ struct AffineLoadOpInterfaceReverse
}
} else {
bool hasIndex = loadOp.getAffineMap().getNumDims() > 0;
auto alignAttr = loadOp->getAttrOfType<IntegerAttr>("alignment");
// if index had to be cached, the pop is not necessarily a valid index
if (hasIndex) {
SmallVector<Value> 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);
}
}
}
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, loadOp.getAlignmentAttr());
}
}
}
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