From 2b5f29c519716b61fcdfec35170f28a8d243883c Mon Sep 17 00:00:00 2001 From: "William S. Moses" Date: Thu, 8 Oct 2026 12:27:19 -0500 Subject: [PATCH] AutoBatching: batch through a nested loop whose trip count the data gives The parallel-while batcher ran a nested loop over the batched values only where its trip count was a constant, though all it needs is that the count be the same for every iteration of the parallel loop, which its condition check already demands: a CSR transpose padded to its longest row runs its inner loop to a max read from the offsets. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_016zErYp7upmqr4NHfhod9UD --- src/enzyme_ad/jax/Passes/AutoBatching.cpp | 11 +- .../parallel_while_invariant_trip.mlir | 105 ++++++++++++++++++ 2 files changed, 111 insertions(+), 5 deletions(-) create mode 100644 test/lit_tests/parallel_while_invariant_trip.mlir diff --git a/src/enzyme_ad/jax/Passes/AutoBatching.cpp b/src/enzyme_ad/jax/Passes/AutoBatching.cpp index 663ea156de..f137539a09 100644 --- a/src/enzyme_ad/jax/Passes/AutoBatching.cpp +++ b/src/enzyme_ad/jax/Passes/AutoBatching.cpp @@ -3043,9 +3043,10 @@ namespace { // Batches the body of an enzymexla.parallel while over its iterations: a value // that varies with the iteration gains a leading dimension of the trip count, // a write into a carried buffer becomes one scatter of every iteration's -// write, and a constant-trip loop nested in the body keeps running, over the -// batched values (the iterations of the parallel loop are independent, so it -// and the parallel loop interchange). +// write, and a loop nested in the body whose trip count is the same for +// every iteration keeps running, over the batched values (the iterations of +// the parallel loop are independent, so it and the parallel loop +// interchange). struct ParallelWhileBatcher { PatternRewriter &rewriter; Location loc; @@ -3165,8 +3166,8 @@ struct ParallelWhileBatcher { // iteration of the parallel loop, over values that may vary with it. LogicalResult analyzeWhile(stablehlo::WhileOp w) { enzyme::WhileLoopInfo wi(w); - if (failed(wi.computeInfo()) || !wi.isValid() || !wi.isConstant() || - !isMemoryEffectFree(w)) + if (failed(wi.computeInfo()) || !wi.isValid() || !wi.isConstantStart() || + !wi.isConstantStep() || !isMemoryEffectFree(w)) return failure(); Value wiv = wi.getInductionVariable(); if (!wiv) diff --git a/test/lit_tests/parallel_while_invariant_trip.mlir b/test/lit_tests/parallel_while_invariant_trip.mlir new file mode 100644 index 0000000000..3acf2a0bcc --- /dev/null +++ b/test/lit_tests/parallel_while_invariant_trip.mlir @@ -0,0 +1,105 @@ +// RUN: enzymexlamlir-opt %s --enzyme-hlo-generate-td="patterns=parallel_while_to_batched_scatter" --transform-interpreter --enzyme-hlo-remove-transform | FileCheck %s + +// A CSR transpose padded to its longest row: the nested loop runs to a +// trip count read from the data, the same for every row, so it keeps +// running over the batched rows. +// M = max_i off[i+1] - off[i]; for i: acc = 0; for k < M: j = off[i] + k; +// if j < off[i+1]: acc += x[idx[j]]; y[i] = acc +func.func @csr_padded(%off: tensor<9xi32>, %idx: tensor<40xi32>, %x: tensor<40xf64>, %y: tensor<8xf64>) -> tensor<8xf64> { + %c0 = stablehlo.constant dense<0> : tensor + %c1 = stablehlo.constant dense<1> : tensor + %c8 = stablehlo.constant dense<8> : tensor + %min = stablehlo.constant dense<-2147483648> : tensor + %zero = stablehlo.constant dense<0.0> : tensor + %lo = stablehlo.slice %off [0:8] : (tensor<9xi32>) -> tensor<8xi32> + %hi = stablehlo.slice %off [1:9] : (tensor<9xi32>) -> tensor<8xi32> + %len = stablehlo.subtract %hi, %lo : tensor<8xi32> + %m = stablehlo.reduce(%len init: %min) applies stablehlo.maximum across dimensions = [0] : (tensor<8xi32>, tensor) -> tensor + %m64 = stablehlo.convert %m : (tensor) -> tensor + %0:2 = stablehlo.while(%i = %c0, %buf = %y) : tensor, tensor<8xf64> attributes {enzymexla.parallel} + cond { + %c = stablehlo.compare LT, %i, %c8 : (tensor, tensor) -> tensor + stablehlo.return %c : tensor + } do { + %l = stablehlo.dynamic_slice %off, %i, sizes = [1] : (tensor<9xi32>, tensor) -> tensor<1xi32> + %ls = stablehlo.reshape %l : (tensor<1xi32>) -> tensor + %i1 = stablehlo.add %i, %c1 : tensor + %h = stablehlo.dynamic_slice %off, %i1, sizes = [1] : (tensor<9xi32>, tensor) -> tensor<1xi32> + %hs = stablehlo.reshape %h : (tensor<1xi32>) -> tensor + %1:2 = stablehlo.while(%k = %c0, %acc = %zero) : tensor, tensor + cond { + %c = stablehlo.compare LT, %k, %m64 : (tensor, tensor) -> tensor + stablehlo.return %c : tensor + } do { + %k32 = stablehlo.convert %k : (tensor) -> tensor + %j = stablehlo.add %ls, %k32 : tensor + %in = stablehlo.compare LT, %j, %hs, SIGNED : (tensor, tensor) -> tensor + %j64 = stablehlo.convert %j : (tensor) -> tensor + %jr = stablehlo.reshape %j64 : (tensor) -> tensor<1xi64> + %col = "stablehlo.gather"(%idx, %jr) <{dimension_numbers = #stablehlo.gather, indices_are_sorted = false, slice_sizes = array}> : (tensor<40xi32>, tensor<1xi64>) -> tensor + %col64 = stablehlo.convert %col : (tensor) -> tensor + %cr = stablehlo.reshape %col64 : (tensor) -> tensor<1xi64> + %v = "stablehlo.gather"(%x, %cr) <{dimension_numbers = #stablehlo.gather, indices_are_sorted = false, slice_sizes = array}> : (tensor<40xf64>, tensor<1xi64>) -> tensor + %vm = stablehlo.select %in, %v, %zero : tensor, tensor + %acc2 = stablehlo.add %acc, %vm : tensor + %kn = stablehlo.add %k, %c1 : tensor + stablehlo.return %kn, %acc2 : tensor, tensor + } + %ir = stablehlo.reshape %i : (tensor) -> tensor<1xi64> + %b2 = "stablehlo.scatter"(%buf, %ir, %1#1) <{indices_are_sorted = false, scatter_dimension_numbers = #stablehlo.scatter, unique_indices = true}> ({ + ^bb0(%p: tensor, %q: tensor): + stablehlo.return %q : tensor + }) : (tensor<8xf64>, tensor<1xi64>, tensor) -> tensor<8xf64> + %in2 = stablehlo.add %i, %c1 : tensor + stablehlo.return %in2, %b2 : tensor, tensor<8xf64> + } + return %0#1 : tensor<8xf64> +} + +// CHECK: func.func @csr_padded(%arg0: tensor<9xi32>, %arg1: tensor<40xi32>, %arg2: tensor<40xf64>, %arg3: tensor<8xf64>) -> tensor<8xf64> { +// CHECK-NEXT: %c = stablehlo.constant dense<0> : tensor +// CHECK-NEXT: %c_0 = stablehlo.constant dense<1> : tensor +// CHECK-NEXT: %c_1 = stablehlo.constant dense<-2147483648> : tensor +// CHECK-NEXT: %cst = stablehlo.constant dense<0.000000e+00> : tensor +// CHECK-NEXT: %0 = stablehlo.slice %arg0 [0:8] : (tensor<9xi32>) -> tensor<8xi32> +// CHECK-NEXT: %1 = stablehlo.slice %arg0 [1:9] : (tensor<9xi32>) -> tensor<8xi32> +// CHECK-NEXT: %2 = stablehlo.subtract %1, %0 : tensor<8xi32> +// CHECK-NEXT: %3 = stablehlo.reduce(%2 init: %c_1) applies stablehlo.maximum across dimensions = [0] : (tensor<8xi32>, tensor) -> tensor +// CHECK-NEXT: %4 = stablehlo.convert %3 : (tensor) -> tensor +// CHECK-NEXT: %5 = stablehlo.iota dim = 0 : tensor<8xi64> +// CHECK-NEXT: %6 = stablehlo.slice %arg0 [0:8] : (tensor<9xi32>) -> tensor<8xi32> +// CHECK-NEXT: %7 = stablehlo.reshape %6 : (tensor<8xi32>) -> tensor<8x1xi32> +// CHECK-NEXT: %8 = stablehlo.reshape %7 : (tensor<8x1xi32>) -> tensor<8xi32> +// CHECK-NEXT: %9 = stablehlo.slice %arg0 [1:9] : (tensor<9xi32>) -> tensor<8xi32> +// CHECK-NEXT: %10 = stablehlo.reshape %9 : (tensor<8xi32>) -> tensor<8x1xi32> +// CHECK-NEXT: %11 = stablehlo.reshape %10 : (tensor<8x1xi32>) -> tensor<8xi32> +// CHECK-NEXT: %12 = stablehlo.broadcast_in_dim %cst, dims = [] : (tensor) -> tensor<8xf64> +// CHECK-NEXT: %13:2 = stablehlo.while(%iterArg = %c, %iterArg_2 = %12) : tensor, tensor<8xf64> +// CHECK-NEXT: cond { +// CHECK-NEXT: %16 = stablehlo.compare LT, %iterArg, %4 : (tensor, tensor) -> tensor +// CHECK-NEXT: stablehlo.return %16 : tensor +// CHECK-NEXT: } do { +// CHECK-NEXT: %16 = stablehlo.convert %iterArg : (tensor) -> tensor +// CHECK-NEXT: %17 = stablehlo.broadcast_in_dim %16, dims = [] : (tensor) -> tensor<8xi32> +// CHECK-NEXT: %18 = stablehlo.add %8, %17 : tensor<8xi32> +// CHECK-NEXT: %19 = stablehlo.compare LT, %18, %11, SIGNED : (tensor<8xi32>, tensor<8xi32>) -> tensor<8xi1> +// CHECK-NEXT: %20 = stablehlo.convert %18 : (tensor<8xi32>) -> tensor<8xi64> +// CHECK-NEXT: %21 = stablehlo.reshape %20 : (tensor<8xi64>) -> tensor<8x1xi64> +// CHECK-NEXT: %22 = "stablehlo.gather"(%arg1, %21) <{dimension_numbers = #stablehlo.gather, indices_are_sorted = false, slice_sizes = array}> : (tensor<40xi32>, tensor<8x1xi64>) -> tensor<8xi32> +// CHECK-NEXT: %23 = stablehlo.convert %22 : (tensor<8xi32>) -> tensor<8xi64> +// CHECK-NEXT: %24 = stablehlo.reshape %23 : (tensor<8xi64>) -> tensor<8x1xi64> +// CHECK-NEXT: %25 = "stablehlo.gather"(%arg2, %24) <{dimension_numbers = #stablehlo.gather, indices_are_sorted = false, slice_sizes = array}> : (tensor<40xf64>, tensor<8x1xi64>) -> tensor<8xf64> +// CHECK-NEXT: %26 = stablehlo.broadcast_in_dim %cst, dims = [] : (tensor) -> tensor<8xf64> +// CHECK-NEXT: %27 = stablehlo.broadcast_in_dim %19, dims = [0] : (tensor<8xi1>) -> tensor<8xi1> +// CHECK-NEXT: %28 = stablehlo.select %27, %25, %26 : tensor<8xi1>, tensor<8xf64> +// CHECK-NEXT: %29 = stablehlo.add %iterArg_2, %28 : tensor<8xf64> +// CHECK-NEXT: %30 = stablehlo.add %iterArg, %c_0 : tensor +// CHECK-NEXT: stablehlo.return %30, %29 : tensor, tensor<8xf64> +// CHECK-NEXT: } +// CHECK-NEXT: %14 = stablehlo.reshape %5 : (tensor<8xi64>) -> tensor<8x1xi64> +// CHECK-NEXT: %15 = "stablehlo.scatter"(%arg3, %14, %13#1) <{indices_are_sorted = false, scatter_dimension_numbers = #stablehlo.scatter, unique_indices = false}> ({ +// CHECK-NEXT: ^bb0(%arg4: tensor, %arg5: tensor): +// CHECK-NEXT: stablehlo.return %arg5 : tensor +// CHECK-NEXT: }) : (tensor<8xf64>, tensor<8x1xi64>, tensor<8xf64>) -> tensor<8xf64> +// CHECK-NEXT: return %15 : tensor<8xf64> +// CHECK-NEXT: }