diff --git a/src/enzyme_ad/jax/Passes/AffineCFG.cpp b/src/enzyme_ad/jax/Passes/AffineCFG.cpp index 3ed0e302ba..e223e0deae 100644 --- a/src/enzyme_ad/jax/Passes/AffineCFG.cpp +++ b/src/enzyme_ad/jax/Passes/AffineCFG.cpp @@ -8843,7 +8843,12 @@ 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) { @@ -8851,9 +8856,18 @@ static bool isLoopMemoryLockStepExecutable(AffineForOp forOp) { << "src: " << *srcOp << "\n" << "dst: " << *dstOp << "\n"); MemRefAccess dstAccess(dstOp); - SmallVector dcs; - DependenceResult result = checkMemrefAccessDependence( - srcAccess, dstAccess, depth, nullptr, &dcs); + DependenceResult result(DependenceResult::NoDependence); + if (srcAccess.memref == dstAccess.memref && + (isa(srcOp) || + isa(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"); diff --git a/test/lit_tests/raising/affine_to_stablehlo_if_masking.mlir b/test/lit_tests/raising/affine_to_stablehlo_if_masking.mlir index 8649a81545..ed5f2f7588 100644 --- a/test/lit_tests/raising/affine_to_stablehlo_if_masking.mlir +++ b/test/lit_tests/raising/affine_to_stablehlo_if_masking.mlir @@ -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 -// CHECK-NEXT: %[[a5]] = stablehlo.constant dense<10> : tensor -// CHECK-NEXT: %[[a7]] = stablehlo.constant dense<1> : tensor -// CHECK-NEXT: %[[a2]]:3 = stablehlo.while(%[[a23:.+]] = %[[a3]], %[[a24:.+]] = %[[a1]], %[[a25:.+]] = %[[a22]]) : tensor, tensor<10xf32>, tensor<10xi1> -// CHECK-NEXT: cond { -// CHECK-NEXT: %[[a4]] = stablehlo.compare LT, %[[a23]], %[[a5]] : (tensor, tensor) -> tensor -// CHECK-NEXT: stablehlo.return %[[a4]] : tensor -// CHECK-NEXT: } do { -// CHECK-NEXT: %[[a4]] = stablehlo.dynamic_slice %[[a25]], %[[a23]], sizes = [1] : (tensor<10xi1>, tensor) -> tensor<1xi1> -// CHECK-NEXT: %[[a6]] = stablehlo.reshape %[[a4]] : (tensor<1xi1>) -> tensor -// CHECK-NEXT: %[[a8]] = stablehlo.dynamic_slice %[[a24]], %[[a23]], sizes = [1] : (tensor<10xf32>, tensor) -> tensor<1xf32> -// CHECK-NEXT: %[[a9]] = stablehlo.reshape %[[a8]] : (tensor<1xf32>) -> tensor -// CHECK-NEXT: %[[a11]] = arith.addf %[[a9]], %[[a9]] : tensor -// 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) -> tensor<1xf32> -// CHECK-NEXT: %[[a20]] = stablehlo.dynamic_slice %[[a24]], %[[a23]], sizes = [1] : (tensor<10xf32>, tensor) -> tensor<1xf32> -// CHECK-NEXT: %[[a21]] = stablehlo.broadcast_in_dim %[[a6]], dims = [] : (tensor) -> 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) -> tensor<10xf32> -// CHECK-NEXT: %[[a28:.+]] = stablehlo.not %[[a6]] : tensor -// CHECK-NEXT: %[[a29:.+]] = stablehlo.dynamic_slice %[[a27]], %[[a23]], sizes = [1] : (tensor<10xf32>, tensor) -> tensor<1xf32> -// CHECK-NEXT: %[[a30:.+]] = stablehlo.reshape %[[a29]] : (tensor<1xf32>) -> tensor -// CHECK-NEXT: %[[a31:.+]] = arith.mulf %[[a30]], %[[a30]] : tensor -// 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) -> tensor<1xf32> -// CHECK-NEXT: %[[a37:.+]] = stablehlo.dynamic_slice %[[a27]], %[[a23]], sizes = [1] : (tensor<10xf32>, tensor) -> tensor<1xf32> -// CHECK-NEXT: %[[a38:.+]] = stablehlo.broadcast_in_dim %[[a28]], dims = [] : (tensor) -> 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) -> tensor<10xf32> -// CHECK-NEXT: %[[a41:.+]] = stablehlo.add %[[a23]], %[[a7]] : tensor -// CHECK-NEXT: stablehlo.return %[[a41]], %[[a40]], %[[a25]] : tensor, 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 +// CHECK-NEXT: %c_2 = stablehlo.constant dense<0> : tensor +// 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 +// 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) -> tensor<10xf32> +// CHECK-NEXT: %6 = stablehlo.not %arg1 : tensor<10xi1> +// CHECK-NEXT: %c_9 = stablehlo.constant dense<0> : tensor +// 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 +// 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) -> 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> diff --git a/test/lit_tests/raising/affine_to_stablehlo_lockstep_invariant_terms.mlir b/test/lit_tests/raising/affine_to_stablehlo_lockstep_invariant_terms.mlir new file mode 100644 index 0000000000..48f94f6763 --- /dev/null +++ b/test/lit_tests/raising/affine_to_stablehlo_lockstep_invariant_terms.mlir @@ -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 +// CHECK-NEXT: %2 = stablehlo.slice %arg2 [1:2] : (tensor<2xi64>) -> tensor<1xi64> +// CHECK-NEXT: %3 = stablehlo.reshape %2 : (tensor<1xi64>) -> tensor +// 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 +// CHECK-NEXT: %10 = stablehlo.broadcast_in_dim %c_3, dims = [] : (tensor) -> 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 +// CHECK-NEXT: %16 = stablehlo.broadcast_in_dim %15, dims = [] : (tensor) -> 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, indices_are_sorted = false, slice_sizes = array}> : (tensor<64xf64>, tensor<3x4x1xi64>) -> tensor<3x4xf64> +// CHECK-NEXT: %c_4 = stablehlo.constant dense<3> : tensor +// CHECK-NEXT: %20 = stablehlo.broadcast_in_dim %c_4, dims = [] : (tensor) -> 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 +// CHECK-NEXT: %26 = stablehlo.broadcast_in_dim %25, dims = [] : (tensor) -> 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, unique_indices = true}> ({ +// CHECK-NEXT: ^bb0(%arg3: tensor, %arg4: tensor): +// CHECK-NEXT: stablehlo.return %arg4 : tensor +// CHECK-NEXT: }) : (tensor<64xf64>, tensor<3x4x1xi64>, tensor<3x4xf64>) -> tensor<64xf64> +// CHECK-NEXT: return %arg0, %29, %arg2 : tensor<64xf64>, tensor<64xf64>, tensor<2xi64> +// CHECK-NEXT: }