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
22 changes: 18 additions & 4 deletions src/enzyme_ad/jax/Passes/AffineCFG.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -8843,17 +8843,31 @@ static bool isLoopMemoryLockStepExecutable(AffineForOp forOp) {
// Dep check depth would be number of enclosing loops + 1.
unsigned depth = ::getNestingDepth(forOp) + 1;

// Check dependences between all pairs of ops in 'loadAndStoreOps'.
// Check dependences between all pairs of ops in 'loadAndStoreOps', on the
// access relations isLoopMemoryParallel builds (loop-invariant terms
// abstracted, with what the nest's in-bounds accesses and the checks on the
// way to it establish). Accesses to different memrefs, or that only read,
// have none.
InvariantTerms terms(forOp);
for (auto *srcOp : loadAndStoreOps) {
MemRefAccess srcAccess(srcOp);
for (auto *dstOp : loadAndStoreOps) {
LLVM_DEBUG(llvm::dbgs() << "Checking dep\n"
<< "src: " << *srcOp << "\n"
<< "dst: " << *dstOp << "\n");
MemRefAccess dstAccess(dstOp);
SmallVector<DependenceComponent, 2> dcs;
DependenceResult result = checkMemrefAccessDependence(
srcAccess, dstAccess, depth, nullptr, &dcs);
DependenceResult result(DependenceResult::NoDependence);
if (srcAccess.memref == dstAccess.memref &&
(isa<AffineWriteOpInterface>(srcOp) ||
isa<AffineWriteOpInterface>(dstOp))) {
presburger::IntegerRelation srcRel(
presburger::PresburgerSpace::getRelationSpace()),
dstRel(presburger::PresburgerSpace::getRelationSpace());
if (failed(terms.accessRelation(srcOp, srcRel)) ||
failed(terms.accessRelation(dstOp, dstRel)))
return false;
result = checkAccessDependence(srcRel, dstRel, depth);
}

if (result.value == DependenceResult::Failure) {
LLVM_DEBUG(llvm::dbgs() << "Failed\n");
Expand Down
69 changes: 28 additions & 41 deletions test/lit_tests/raising/affine_to_stablehlo_if_masking.mlir
Original file line number Diff line number Diff line change
Expand Up @@ -78,47 +78,34 @@ func.func @test_affine_if_masking(%arg0: memref<10xf32>) {
// CHECK-NEXT: return %[[a21]] : tensor<10xf32>
// CHECK-NEXT: }
// CHECK-NEXT: func.func private @test_if_else_masking_raised(%[[a1]]: tensor<10xf32>, %[[a22:.+]]: tensor<10xi1>) -> (tensor<10xf32>, tensor<10xi1>) {
// CHECK-NEXT: %[[a3]] = stablehlo.constant dense<0> : tensor<i64>
// CHECK-NEXT: %[[a5]] = stablehlo.constant dense<10> : tensor<i64>
// CHECK-NEXT: %[[a7]] = stablehlo.constant dense<1> : tensor<i64>
// CHECK-NEXT: %[[a2]]:3 = stablehlo.while(%[[a23:.+]] = %[[a3]], %[[a24:.+]] = %[[a1]], %[[a25:.+]] = %[[a22]]) : tensor<i64>, tensor<10xf32>, tensor<10xi1>
// CHECK-NEXT: cond {
// CHECK-NEXT: %[[a4]] = stablehlo.compare LT, %[[a23]], %[[a5]] : (tensor<i64>, tensor<i64>) -> tensor<i1>
// CHECK-NEXT: stablehlo.return %[[a4]] : tensor<i1>
// CHECK-NEXT: } do {
// CHECK-NEXT: %[[a4]] = stablehlo.dynamic_slice %[[a25]], %[[a23]], sizes = [1] : (tensor<10xi1>, tensor<i64>) -> tensor<1xi1>
// CHECK-NEXT: %[[a6]] = stablehlo.reshape %[[a4]] : (tensor<1xi1>) -> tensor<i1>
// CHECK-NEXT: %[[a8]] = stablehlo.dynamic_slice %[[a24]], %[[a23]], sizes = [1] : (tensor<10xf32>, tensor<i64>) -> tensor<1xf32>
// CHECK-NEXT: %[[a9]] = stablehlo.reshape %[[a8]] : (tensor<1xf32>) -> tensor<f32>
// CHECK-NEXT: %[[a11]] = arith.addf %[[a9]], %[[a9]] : tensor<f32>
// CHECK-NEXT: %[[a14]] = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %[[a15]] = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %[[a16]] = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %[[a17]] = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %[[a18]] = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %[[a13]] = stablehlo.broadcast_in_dim %[[a11]], dims = [] : (tensor<f32>) -> tensor<1xf32>
// CHECK-NEXT: %[[a20]] = stablehlo.dynamic_slice %[[a24]], %[[a23]], sizes = [1] : (tensor<10xf32>, tensor<i64>) -> tensor<1xf32>
// CHECK-NEXT: %[[a21]] = stablehlo.broadcast_in_dim %[[a6]], dims = [] : (tensor<i1>) -> tensor<1xi1>
// CHECK-NEXT: %[[a26:.+]] = stablehlo.select %[[a21]], %[[a13]], %[[a20]] : tensor<1xi1>, tensor<1xf32>
// CHECK-NEXT: %[[a27:.+]] = stablehlo.dynamic_update_slice %[[a24]], %[[a26]], %[[a23]] : (tensor<10xf32>, tensor<1xf32>, tensor<i64>) -> tensor<10xf32>
// CHECK-NEXT: %[[a28:.+]] = stablehlo.not %[[a6]] : tensor<i1>
// CHECK-NEXT: %[[a29:.+]] = stablehlo.dynamic_slice %[[a27]], %[[a23]], sizes = [1] : (tensor<10xf32>, tensor<i64>) -> tensor<1xf32>
// CHECK-NEXT: %[[a30:.+]] = stablehlo.reshape %[[a29]] : (tensor<1xf32>) -> tensor<f32>
// CHECK-NEXT: %[[a31:.+]] = arith.mulf %[[a30]], %[[a30]] : tensor<f32>
// CHECK-NEXT: %[[a19]] = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %[[a32:.+]] = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %[[a33:.+]] = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %[[a34:.+]] = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %[[a35:.+]] = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %[[a36:.+]] = stablehlo.broadcast_in_dim %[[a31]], dims = [] : (tensor<f32>) -> tensor<1xf32>
// CHECK-NEXT: %[[a37:.+]] = stablehlo.dynamic_slice %[[a27]], %[[a23]], sizes = [1] : (tensor<10xf32>, tensor<i64>) -> tensor<1xf32>
// CHECK-NEXT: %[[a38:.+]] = stablehlo.broadcast_in_dim %[[a28]], dims = [] : (tensor<i1>) -> tensor<1xi1>
// CHECK-NEXT: %[[a39:.+]] = stablehlo.select %[[a38]], %[[a36]], %[[a37]] : tensor<1xi1>, tensor<1xf32>
// CHECK-NEXT: %[[a40:.+]] = stablehlo.dynamic_update_slice %[[a27]], %[[a39]], %[[a23]] : (tensor<10xf32>, tensor<1xf32>, tensor<i64>) -> tensor<10xf32>
// CHECK-NEXT: %[[a41:.+]] = stablehlo.add %[[a23]], %[[a7]] : tensor<i64>
// CHECK-NEXT: stablehlo.return %[[a41]], %[[a40]], %[[a25]] : tensor<i64>, tensor<10xf32>, tensor<10xi1>
// CHECK-NEXT: }
// CHECK-NEXT: return %[[a2]]#1, %[[a2]]#2 : tensor<10xf32>, tensor<10xi1>
// CHECK-NEXT: %0 = stablehlo.iota dim = 0 : tensor<10xi64>
// CHECK-NEXT: %c = stablehlo.constant dense<0> : tensor<10xi64>
// CHECK-NEXT: %1 = stablehlo.add %0, %c : tensor<10xi64>
// CHECK-NEXT: %c_0 = stablehlo.constant dense<1> : tensor<10xi64>
// CHECK-NEXT: %2 = stablehlo.multiply %1, %c_0 : tensor<10xi64>
// CHECK-NEXT: %c_1 = stablehlo.constant dense<0> : tensor<i64>
// CHECK-NEXT: %c_2 = stablehlo.constant dense<0> : tensor<i64>
// CHECK-NEXT: %3 = arith.addf %arg0, %arg0 : tensor<10xf32>
// CHECK-NEXT: %c_3 = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %c_4 = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %c_5 = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %c_6 = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %c_7 = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %c_8 = stablehlo.constant dense<0> : tensor<i64>
// CHECK-NEXT: %4 = stablehlo.select %arg1, %3, %arg0 : tensor<10xi1>, tensor<10xf32>
// CHECK-NEXT: %5 = stablehlo.dynamic_update_slice %arg0, %4, %c_8 : (tensor<10xf32>, tensor<10xf32>, tensor<i64>) -> tensor<10xf32>
// CHECK-NEXT: %6 = stablehlo.not %arg1 : tensor<10xi1>
// CHECK-NEXT: %c_9 = stablehlo.constant dense<0> : tensor<i64>
// CHECK-NEXT: %7 = arith.mulf %5, %5 : tensor<10xf32>
// CHECK-NEXT: %c_10 = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %c_11 = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %c_12 = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %c_13 = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %c_14 = stablehlo.constant dense<0> : tensor<1xi64>
// CHECK-NEXT: %c_15 = stablehlo.constant dense<0> : tensor<i64>
// CHECK-NEXT: %8 = stablehlo.select %6, %7, %5 : tensor<10xi1>, tensor<10xf32>
// CHECK-NEXT: %9 = stablehlo.dynamic_update_slice %5, %8, %c_15 : (tensor<10xf32>, tensor<10xf32>, tensor<i64>) -> tensor<10xf32>
// CHECK-NEXT: return %9, %arg1 : tensor<10xf32>, tensor<10xi1>
// CHECK-NEXT: }
// CHECK-NEXT: func.func private @test_nested_if_masking_raised(%[[a1]]: tensor<10xf32>, %[[a22]]: tensor<10xi1>, %[[a42:.+]]: tensor<10xi1>) -> (tensor<10xf32>, tensor<10xi1>, tensor<10xi1>) {
// CHECK-NEXT: %[[a2]] = stablehlo.iota dim = 0 : tensor<10xi64>
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,64 @@
// RUN: enzymexlamlir-opt %s "--pass-pipeline=builtin.module(raise-affine-to-stablehlo{enable_lockstep_for=true})" | FileCheck %s

// A loop of constant count whose accesses are offset by a product of
// symbols: the lockstep check now builds its dependence problems as the
// parallel check does, with that product abstracted, and the loop runs in
// lock step (one gather and one scatter) rather than as a while.
module {
func.func private @copy(%a: memref<64xf64, 1>, %b: memref<64xf64, 1>, %p: memref<2xi64, 1>) {
%s64 = affine.load %p[0] : memref<2xi64, 1>
%t64 = affine.load %p[1] : memref<2xi64, 1>
%s = arith.index_cast %s64 : i64 to index
%t = arith.index_cast %t64 : i64 to index
affine.parallel (%i) = (0) to (4) {
affine.for %k = 0 to 3 {
%v = affine.load %a[%k + %i * 3 + symbol(%s) * symbol(%t)] : memref<64xf64, 1>
affine.store %v, %b[%k + %i * 3 + symbol(%s) * symbol(%t)] : memref<64xf64, 1>
}
}
return
}
}

// CHECK: func.func private @copy_raised(%arg0: tensor<64xf64>, %arg1: tensor<64xf64>, %arg2: tensor<2xi64>) -> (tensor<64xf64>, tensor<64xf64>, tensor<2xi64>) {
// CHECK-NEXT: %0 = stablehlo.slice %arg2 [0:1] : (tensor<2xi64>) -> tensor<1xi64>
// CHECK-NEXT: %1 = stablehlo.reshape %0 : (tensor<1xi64>) -> tensor<i64>
// CHECK-NEXT: %2 = stablehlo.slice %arg2 [1:2] : (tensor<2xi64>) -> tensor<1xi64>
// CHECK-NEXT: %3 = stablehlo.reshape %2 : (tensor<1xi64>) -> tensor<i64>
// CHECK-NEXT: %4 = stablehlo.iota dim = 0 : tensor<4xi64>
// CHECK-NEXT: %c = stablehlo.constant dense<0> : tensor<4xi64>
// CHECK-NEXT: %5 = stablehlo.add %4, %c : tensor<4xi64>
// CHECK-NEXT: %c_0 = stablehlo.constant dense<1> : tensor<4xi64>
// CHECK-NEXT: %6 = stablehlo.multiply %5, %c_0 : tensor<4xi64>
// CHECK-NEXT: %7 = stablehlo.iota dim = 0 : tensor<3xi64>
// CHECK-NEXT: %c_1 = stablehlo.constant dense<0> : tensor<3xi64>
// CHECK-NEXT: %8 = stablehlo.add %7, %c_1 : tensor<3xi64>
// CHECK-NEXT: %c_2 = stablehlo.constant dense<1> : tensor<3xi64>
// CHECK-NEXT: %9 = stablehlo.multiply %8, %c_2 : tensor<3xi64>
// CHECK-NEXT: %c_3 = stablehlo.constant dense<3> : tensor<i64>
// CHECK-NEXT: %10 = stablehlo.broadcast_in_dim %c_3, dims = [] : (tensor<i64>) -> tensor<4xi64>
// CHECK-NEXT: %11 = stablehlo.multiply %6, %10 : tensor<4xi64>
// CHECK-NEXT: %12 = stablehlo.broadcast_in_dim %9, dims = [0] : (tensor<3xi64>) -> tensor<3x4xi64>
// CHECK-NEXT: %13 = stablehlo.broadcast_in_dim %11, dims = [1] : (tensor<4xi64>) -> tensor<3x4xi64>
// CHECK-NEXT: %14 = stablehlo.add %12, %13 : tensor<3x4xi64>
// CHECK-NEXT: %15 = stablehlo.multiply %1, %3 : tensor<i64>
// CHECK-NEXT: %16 = stablehlo.broadcast_in_dim %15, dims = [] : (tensor<i64>) -> tensor<3x4xi64>
// CHECK-NEXT: %17 = stablehlo.add %14, %16 : tensor<3x4xi64>
// CHECK-NEXT: %18 = stablehlo.reshape %17 : (tensor<3x4xi64>) -> tensor<3x4x1xi64>
// CHECK-NEXT: %19 = "stablehlo.gather"(%arg0, %18) <{dimension_numbers = #stablehlo.gather<collapsed_slice_dims = [0], start_index_map = [0], index_vector_dim = 2>, indices_are_sorted = false, slice_sizes = array<i64: 1>}> : (tensor<64xf64>, tensor<3x4x1xi64>) -> tensor<3x4xf64>
// CHECK-NEXT: %c_4 = stablehlo.constant dense<3> : tensor<i64>
// CHECK-NEXT: %20 = stablehlo.broadcast_in_dim %c_4, dims = [] : (tensor<i64>) -> tensor<4xi64>
// CHECK-NEXT: %21 = stablehlo.multiply %6, %20 : tensor<4xi64>
// CHECK-NEXT: %22 = stablehlo.broadcast_in_dim %9, dims = [0] : (tensor<3xi64>) -> tensor<3x4xi64>
// CHECK-NEXT: %23 = stablehlo.broadcast_in_dim %21, dims = [1] : (tensor<4xi64>) -> tensor<3x4xi64>
// CHECK-NEXT: %24 = stablehlo.add %22, %23 : tensor<3x4xi64>
// CHECK-NEXT: %25 = stablehlo.multiply %1, %3 : tensor<i64>
// CHECK-NEXT: %26 = stablehlo.broadcast_in_dim %25, dims = [] : (tensor<i64>) -> tensor<3x4xi64>
// CHECK-NEXT: %27 = stablehlo.add %24, %26 : tensor<3x4xi64>
// CHECK-NEXT: %28 = stablehlo.reshape %27 : (tensor<3x4xi64>) -> tensor<3x4x1xi64>
// CHECK-NEXT: %29 = "stablehlo.scatter"(%arg1, %28, %19) <{indices_are_sorted = false, scatter_dimension_numbers = #stablehlo.scatter<inserted_window_dims = [0], scatter_dims_to_operand_dims = [0], index_vector_dim = 2>, unique_indices = true}> ({
// CHECK-NEXT: ^bb0(%arg3: tensor<f64>, %arg4: tensor<f64>):
// CHECK-NEXT: stablehlo.return %arg4 : tensor<f64>
// CHECK-NEXT: }) : (tensor<64xf64>, tensor<3x4x1xi64>, tensor<3x4xf64>) -> tensor<64xf64>
// CHECK-NEXT: return %arg0, %29, %arg2 : tensor<64xf64>, tensor<64xf64>, tensor<2xi64>
// CHECK-NEXT: }
Loading