diff --git a/patches/llvm_affine_parallel_signless_minmax.patch b/patches/llvm_affine_parallel_signless_minmax.patch new file mode 100644 index 0000000000..9c4bc1376e --- /dev/null +++ b/patches/llvm_affine_parallel_signless_minmax.patch @@ -0,0 +1,27 @@ +diff --git a/mlir/lib/Dialect/Affine/IR/AffineOps.cpp b/mlir/lib/Dialect/Affine/IR/AffineOps.cpp +--- a/mlir/lib/Dialect/Affine/IR/AffineOps.cpp ++++ b/mlir/lib/Dialect/Affine/IR/AffineOps.cpp +@@ -4281,19 +4281,19 @@ static bool isResultTypeMatchAtomicRMWKind(Type resultType, + return isa(resultType); + case arith::AtomicRMWKind::maxs: { + auto intType = dyn_cast(resultType); +- return intType && intType.isSigned(); ++ return intType && !intType.isUnsigned(); + } + case arith::AtomicRMWKind::mins: { + auto intType = dyn_cast(resultType); +- return intType && intType.isSigned(); ++ return intType && !intType.isUnsigned(); + } + case arith::AtomicRMWKind::maxu: { + auto intType = dyn_cast(resultType); +- return intType && intType.isUnsigned(); ++ return intType && !intType.isSigned(); + } + case arith::AtomicRMWKind::minu: { + auto intType = dyn_cast(resultType); +- return intType && intType.isUnsigned(); ++ return intType && !intType.isSigned(); + } + case arith::AtomicRMWKind::ori: + case arith::AtomicRMWKind::andi: diff --git a/src/enzyme_ad/jax/Passes/AffineCFG.cpp b/src/enzyme_ad/jax/Passes/AffineCFG.cpp index bb5b0b991c..d217ce65b6 100644 --- a/src/enzyme_ad/jax/Passes/AffineCFG.cpp +++ b/src/enzyme_ad/jax/Passes/AffineCFG.cpp @@ -3428,11 +3428,12 @@ struct ForOpRaising : public OpRewritePattern { // The reduction kind whose combining op `op` is (the inverse of // arith::getReductionOp), where affine.parallel admits that kind on op's -// type: its signed and unsigned min/max want an integer of that signedness. +// type: its signed and unsigned min/max want an integer not of the other +// signedness. static std::optional reductionKind(Operation *op) { auto intType = dyn_cast(op->getResult(0).getType()); - bool isSigned = intType && intType.isSigned(); - bool isUnsigned = intType && intType.isUnsigned(); + bool isSigned = intType && !intType.isUnsigned(); + bool isUnsigned = intType && !intType.isSigned(); auto ifType = [](bool ok, AtomicRMWKind kind) -> std::optional { if (ok) @@ -7331,6 +7332,263 @@ struct AffineForCopyCarry : public OpRewritePattern { } }; +// The buffer behind a memref: the pointer a pointer2memref views, or the +// memref itself. +static Value bufferOf(Value memref) { + while (auto p2m = memref.getDefiningOp()) + memref = p2m.getSource(); + return memref; +} + +// Whether something within `loop` may write `memref`'s buffer, or has +// effects that cannot be told. +static bool mayWriteWithin(Operation *loop, Value memref) { + Value buffer = bufferOf(memref); + bool written = false; + loop->walk([&](Operation *op) { + if (written || op->hasTrait()) + return; + auto iface = dyn_cast(op); + if (!iface) { + written = !isMemoryEffectFree(op); + return; + } + SmallVector effects; + iface.getEffects(effects); + for (auto &effect : effects) { + if (isa(effect.getEffect())) + continue; + Value on = effect.getValue(); + if (!on || bufferOf(on) == buffer) + written = true; + } + }); + return written; +} + +// The induction variables of an affine loop. +static ValueRange loopIVs(Operation *loop) { + if (auto forOp = dyn_cast(loop)) + return forOp.getBody()->getArguments().take_front(1); + return cast(loop).getIVs(); +} + +// The rows of the bounds `lo` and `hi` of `loop`: the enclosing affine +// loops whose variables they vary with, outermost first, and the ops that +// compute them (`chain`) within the outermost of those: pure ops, and loads +// of buffers nothing in that loop writes that every iteration of the rows +// runs (under no other op). False where the bounds vary with anything else. +static bool rowsOfBounds(scf::ForOp loop, Value lo, Value hi, Region *scope, + SmallVectorImpl &rows, + SetVector &chain) { + SmallVector todo{lo, hi}; + DenseSet seen; + SetVector rowSet; + SmallVector loads; + while (!todo.empty()) { + Value v = todo.pop_back_val(); + if (!seen.insert(v).second) + continue; + if (auto arg = dyn_cast(v)) { + Operation *owner = arg.getOwner()->getParentOp(); + if (!owner->isProperAncestor(loop)) + return false; + if (isa(owner) && + llvm::is_contained(loopIVs(owner), v)) { + rowSet.insert(owner); + continue; + } + if (isa(owner)) + return false; + continue; + } + Operation *op = v.getDefiningOp(); + if (!scope->isAncestor(op->getParentRegion())) + continue; + if (isa(op)) + loads.push_back(op->getOperand(0)); + else if (!isMemoryEffectFree(op) || op->getNumRegions() != 0) + return false; + chain.insert(op); + llvm::append_range(todo, op->getOperands()); + } + if (rowSet.empty()) + return false; + // outermost first + for (Operation *row : rowSet) + rows.push_back(row); + llvm::sort(rows, + [](Operation *a, Operation *b) { return a->isProperAncestor(b); }); + Operation *outer = rows.front(); + for (Operation *op : chain) { + if (!outer->isAncestor(op) || isSpeculatable(op)) + continue; + for (Operation *parent = op->getParentOp(); parent != outer; + parent = parent->getParentOp()) + if (!rowSet.contains(parent)) + return false; + } + for (Value memref : loads) + if (mayWriteWithin(outer, memref)) + return false; + // The rows' bounds, read before the outermost: from outside it, or the + // variables of enclosing rows. + for (Operation *row : rows) { + SmallVector operands; + if (auto forOp = dyn_cast(row)) { + llvm::append_range(operands, forOp.getLowerBoundOperands()); + llvm::append_range(operands, forOp.getUpperBoundOperands()); + } else { + auto par = cast(row); + llvm::append_range(operands, par.getLowerBoundsOperands()); + llvm::append_range(operands, par.getUpperBoundsOperands()); + } + for (Value v : operands) { + if (!outer->isAncestor(v.getParentRegion()->getParentOp())) + continue; + auto arg = dyn_cast(v); + Operation *owner = arg ? arg.getOwner()->getParentOp() : nullptr; + if (!owner || !rowSet.contains(owner) || !owner->isProperAncestor(row)) + return false; + } + } + return true; +} + +// `v` as the integer type `type`. +static Value asInteger(Value v, Type type, Location loc, + PatternRewriter &rewriter) { + if (v.getType() == type) + return v; + return arith::IndexCastOp::create(rewriter, loc, type, v); +} + +// An scf.for whose bounds a row of a structure gives, as the entries of a +// CSR row, `for (j = I[i]; j < I[i+1]; ++j)`: the bounds are not affine, the +// loop stays an scf.for, and raised it runs the rows one at a time. Run +// instead to the longest row, M = max over the rows of hi - lo, a parallel +// reduction before the row loop, with the body under `lo + k < hi`: every +// row then runs the same k loop, and the rows raise as its lanes. +struct PadLoopToLongestRow : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(scf::ForOp loop, + PatternRewriter &rewriter) const override { + if (loop->hasAttr("enzyme.enable_checkpointing")) + return failure(); + APInt step; + if (!matchPattern(loop.getStep(), m_ConstantInt(&step)) || + !step.isStrictlyPositive()) + return failure(); + Region *scope = getLocalAffineScope(loop); + if (!scope) + return failure(); + Value lo = loop.getLowerBound(), hi = loop.getUpperBound(); + if (isValidIndex(lo, scope) && isValidIndex(hi, scope)) + return failure(); + SmallVector rows; + SetVector chain; + if (!rowsOfBounds(loop, lo, hi, scope, rows, chain)) + return failure(); + Operation *outer = rows.front(); + // M is a symbol only at the top of the scope. + if (!outer->getParentOp()->hasTrait()) + return failure(); + + Location loc = loop.getLoc(); + Type ivType = loop.getInductionVar().getType(); + Type redType = isa(ivType) ? rewriter.getI64Type() : ivType; + + // The reduction: the rows again, yielding the row's length. + rewriter.setInsertionPoint(outer); + IRMapping map; + SmallVector reductions; + for (Operation *row : rows) { + SmallVector lbMaps, ubMaps; + SmallVector lbArgs, ubArgs; + SmallVector steps; + if (auto forOp = dyn_cast(row)) { + lbMaps.push_back(forOp.getLowerBoundMap()); + ubMaps.push_back(forOp.getUpperBoundMap()); + llvm::append_range(lbArgs, forOp.getLowerBoundOperands()); + llvm::append_range(ubArgs, forOp.getUpperBoundOperands()); + steps.push_back(forOp.getStepAsInt()); + } else { + auto par = cast(row); + for (unsigned i = 0, e = par.getNumDims(); i < e; ++i) { + lbMaps.push_back(par.getLowerBoundMap(i)); + ubMaps.push_back(par.getUpperBoundMap(i)); + } + llvm::append_range(lbArgs, par.getLowerBoundsOperands()); + llvm::append_range(ubArgs, par.getUpperBoundsOperands()); + llvm::append_range(steps, par.getSteps()); + } + for (Value &v : lbArgs) + v = map.lookupOrDefault(v); + for (Value &v : ubArgs) + v = map.lookupOrDefault(v); + Type resultTypes[] = {redType}; + arith::AtomicRMWKind kinds[] = {arith::AtomicRMWKind::maxs}; + auto reduction = + AffineParallelOp::create(rewriter, loc, resultTypes, kinds, lbMaps, + lbArgs, ubMaps, ubArgs, steps); + map.map(loopIVs(row), reduction.getIVs()); + rewriter.setInsertionPointToEnd(reduction.getBody()); + reductions.push_back(reduction); + } + // The bounds, computed again in the innermost: their ops within the + // outermost row, in the order they come. + outer->walk([&](Operation *op) { + if (chain.contains(op)) + rewriter.clone(*op, map); + }); + Value length = arith::SubIOp::create(rewriter, loc, map.lookupOrDefault(hi), + map.lookupOrDefault(lo)); + length = asInteger(length, redType, loc, rewriter); + AffineYieldOp::create(rewriter, loc, length); + for (size_t i = reductions.size() - 1; i > 0; --i) { + rewriter.setInsertionPointToEnd(reductions[i - 1].getBody()); + AffineYieldOp::create(rewriter, loc, reductions[i].getResult(0)); + } + rewriter.setInsertionPoint(outer); + Value longest = asInteger(reductions.front().getResult(0), + rewriter.getIndexType(), loc, rewriter); + + // The k loop, its body under `lo + k * step < hi`. + rewriter.setInsertionPoint(loop); + AffineMap ubMap = AffineMap::get(0, 1, rewriter.getAffineSymbolExpr(0)); + auto kLoop = AffineForOp::create(rewriter, loc, ValueRange{}, + rewriter.getConstantAffineMap(0), longest, + ubMap, 1, loop.getInits()); + Block *body = kLoop.getBody(); + if (kLoop.getNumIterOperands() == 0) + rewriter.eraseOp(body->getTerminator()); + rewriter.setInsertionPointToEnd(body); + Value k = asInteger(kLoop.getInductionVar(), ivType, loc, rewriter); + if (step != 1) { + Value stepValue = arith::ConstantOp::create( + rewriter, loc, rewriter.getIntegerAttr(ivType, step)); + k = arith::MulIOp::create(rewriter, loc, k, stepValue); + } + Value j = arith::AddIOp::create(rewriter, loc, lo, k); + Value runs = + arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::slt, j, hi); + auto ifOp = scf::IfOp::create(rewriter, loc, loop.getResultTypes(), runs, + /*addThenBlock=*/true, + /*addElseBlock=*/true); + SmallVector args{j}; + llvm::append_range(args, kLoop.getRegionIterArgs()); + rewriter.inlineBlockBefore(loop.getBody(), ifOp.thenBlock(), + ifOp.thenBlock()->end(), args); + rewriter.setInsertionPointToEnd(ifOp.elseBlock()); + scf::YieldOp::create(rewriter, loc, kLoop.getRegionIterArgs()); + rewriter.setInsertionPointToEnd(body); + AffineYieldOp::create(rewriter, loc, ifOp.getResults()); + rewriter.replaceOp(loop, kLoop.getResults()); + return success(); + } +}; + // The same for an scf.for, as a select on the bounds. struct SCFForCopyCarry : public OpRewritePattern { using OpRewritePattern::OpRewritePattern; @@ -7814,6 +8072,7 @@ void mlir::enzyme::populateAffineCFGPatterns( RewritePatternSet &rpl, bool enable_split_on_affine_if_constants) { MLIRContext *context = rpl.getContext(); mlir::enzyme::addSingleIter(rpl, context); + rpl.add(context, 0); rpl.add, CanonicalizeIndexCast, AffineIfYieldMovementPattern, diff --git a/src/enzyme_ad/jax/Passes/AffineToStableHLORaising.cpp b/src/enzyme_ad/jax/Passes/AffineToStableHLORaising.cpp index 970e4ce199..25fc3153bf 100644 --- a/src/enzyme_ad/jax/Passes/AffineToStableHLORaising.cpp +++ b/src/enzyme_ad/jax/Passes/AffineToStableHLORaising.cpp @@ -3088,6 +3088,43 @@ static LogicalResult tryRaisingParallelOpToStableHLO( builder.getOneAttr(unrankedTensorType)); innerRedName = "stablehlo.multiply"; break; + case arith::AtomicRMWKind::maxs: + inits[0] = stablehlo::ConstantOp::create( + builder, + rewriteLocation(res.getLoc(), pc.options.strip_llvm_debuginfo), + SplatElementsAttr::get(unrankedTensorType, + ArrayRef(IntegerAttr::get( + ET, APInt::getSignedMinValue( + ET.getIntOrFloatBitWidth()))))); + innerRedName = "stablehlo.maximum"; + break; + case arith::AtomicRMWKind::mins: + inits[0] = stablehlo::ConstantOp::create( + builder, + rewriteLocation(res.getLoc(), pc.options.strip_llvm_debuginfo), + SplatElementsAttr::get(unrankedTensorType, + ArrayRef(IntegerAttr::get( + ET, APInt::getSignedMaxValue( + ET.getIntOrFloatBitWidth()))))); + innerRedName = "stablehlo.minimum"; + break; + case arith::AtomicRMWKind::maxu: + inits[0] = stablehlo::ConstantOp::create( + builder, + rewriteLocation(res.getLoc(), pc.options.strip_llvm_debuginfo), + builder.getZeroAttr(unrankedTensorType)); + innerRedName = "arith.maxui"; + break; + case arith::AtomicRMWKind::minu: + inits[0] = stablehlo::ConstantOp::create( + builder, + rewriteLocation(res.getLoc(), pc.options.strip_llvm_debuginfo), + SplatElementsAttr::get( + unrankedTensorType, + ArrayRef(IntegerAttr::get( + ET, APInt::getAllOnes(ET.getIntOrFloatBitWidth()))))); + innerRedName = "arith.minui"; + break; case arith::AtomicRMWKind::ori: inits[0] = stablehlo::ConstantOp::create( builder, diff --git a/test/lit_tests/affinecfg_copy_carry.mlir b/test/lit_tests/affinecfg_copy_carry.mlir index 5ef057c5b5..10edd3c21e 100644 --- a/test/lit_tests/affinecfg_copy_carry.mlir +++ b/test/lit_tests/affinecfg_copy_carry.mlir @@ -80,19 +80,29 @@ func.func @scf_copy(%x: memref, %nb: memref, %m: index, %junk: f } // CHECK: func.func @scf_copy(%arg0: memref, %arg1: memref, %arg2: index, %arg3: f64, %arg4: memref) { -// CHECK-NEXT: %cst = arith.constant 0.000000e+00 : f64 +// CHECK-NEXT: %cst = arith.constant -0.000000e+00 : f64 // CHECK-NEXT: %c0 = arith.constant 0 : index -// CHECK-NEXT: %c1 = arith.constant 1 : index -// CHECK-NEXT: affine.for %arg5 = 0 to %arg2 { -// CHECK-NEXT: %0 = affine.load %arg1[%arg5] : memref -// CHECK-NEXT: %1 = scf.for %arg6 = %c0 to %0 step %c1 iter_args(%arg7 = %cst) -> (f64) { -// CHECK-NEXT: %4 = memref.load %arg0[%arg6] : memref -// CHECK-NEXT: %5 = arith.addf %arg7, %4 : f64 -// CHECK-NEXT: scf.yield %5 : f64 +// CHECK-NEXT: %0 = affine.parallel (%arg5) = (0) to (symbol(%arg2)) reduce ("maxs") -> (i64) { +// CHECK-NEXT: %2 = affine.load %arg1[%arg5] : memref +// CHECK-NEXT: %3 = arith.index_cast %2 : index to i64 +// CHECK-NEXT: affine.yield %3 : i64 +// CHECK-NEXT: } +// CHECK-NEXT: %1 = arith.index_cast %0 : i64 to index +// CHECK-NEXT: affine.parallel (%arg5) = (0) to (symbol(%arg2)) { +// CHECK-NEXT: %2 = affine.load %arg1[%arg5] : memref +// CHECK-NEXT: %3 = affine.parallel (%arg6) = (0) to (symbol(%1)) reduce ("addf") -> (f64) { +// CHECK-NEXT: %6 = arith.cmpi slt, %arg6, %2 : index +// CHECK-NEXT: %7 = scf.if %6 -> (f64) { +// CHECK-NEXT: %8 = affine.load %arg0[%arg6] : memref +// CHECK-NEXT: scf.yield %8 : f64 +// CHECK-NEXT: } else { +// CHECK-NEXT: scf.yield %cst : f64 +// CHECK-NEXT: } +// CHECK-NEXT: affine.yield %7 : f64 // CHECK-NEXT: } -// CHECK-NEXT: %2 = arith.cmpi sgt, %0, %c0 : index -// CHECK-NEXT: %3 = arith.select %2, %1, %arg3 : f64 -// CHECK-NEXT: affine.store %3, %arg4[%arg5] : memref +// CHECK-NEXT: %4 = arith.cmpi sgt, %2, %c0 : index +// CHECK-NEXT: %5 = arith.select %4, %3, %arg3 : f64 +// CHECK-NEXT: affine.store %5, %arg4[%arg5] : memref // CHECK-NEXT: } // CHECK-NEXT: return // CHECK-NEXT: } diff --git a/test/lit_tests/affinecfg_int_max_reduction.mlir b/test/lit_tests/affinecfg_int_max_reduction.mlir new file mode 100644 index 0000000000..7b61578a2d --- /dev/null +++ b/test/lit_tests/affinecfg_int_max_reduction.mlir @@ -0,0 +1,60 @@ +// RUN: enzymexlamlir-opt --affine-cfg --split-input-file %s | FileCheck %s + +// A loop carrying the max of what it reads is a parallel reduction. +func.func @row_max(%off: memref<101xi32>) -> i32 { + %c0_i32 = arith.constant 0 : i32 + %0 = affine.for %i = 0 to 100 iter_args(%m = %c0_i32) -> (i32) { + %lo = affine.load %off[%i] : memref<101xi32> + %hi = affine.load %off[%i + 1] : memref<101xi32> + %d = arith.subi %hi, %lo : i32 + %mx = arith.maxsi %m, %d : i32 + affine.yield %mx : i32 + } + return %0 : i32 +} + +// CHECK: func.func @row_max(%arg0: memref<101xi32>) -> i32 { +// CHECK-NEXT: %c0_i32 = arith.constant 0 : i32 +// CHECK-NEXT: %0 = affine.parallel (%arg1) = (0) to (100) reduce ("maxs") -> (i32) { +// CHECK-NEXT: %2 = affine.load %arg0[%arg1] : memref<101xi32> +// CHECK-NEXT: %3 = affine.load %arg0[%arg1 + 1] : memref<101xi32> +// CHECK-NEXT: %4 = arith.subi %3, %2 : i32 +// CHECK-NEXT: affine.yield %4 : i32 +// CHECK-NEXT: } +// CHECK-NEXT: %1 = arith.maxsi %0, %c0_i32 : i32 +// CHECK-NEXT: return %1 : i32 +// CHECK-NEXT: } + +// ----- + +// The same loop lowered: an scf.parallel reducing with maxsi. +func.func @row_max_scf(%off: memref<101xi32>) -> i32 { + %c0 = arith.constant 0 : index + %c1 = arith.constant 1 : index + %c100 = arith.constant 100 : index + %c0_i32 = arith.constant 0 : i32 + %0 = scf.parallel (%i) = (%c0) to (%c100) step (%c1) init (%c0_i32) -> i32 { + %lo = memref.load %off[%i] : memref<101xi32> + %i1 = arith.addi %i, %c1 : index + %hi = memref.load %off[%i1] : memref<101xi32> + %d = arith.subi %hi, %lo : i32 + scf.reduce(%d : i32) { + ^bb0(%a: i32, %b: i32): + %m = arith.maxsi %a, %b : i32 + scf.reduce.return %m : i32 + } + } + return %0 : i32 +} + +// CHECK: func.func @row_max_scf(%arg0: memref<101xi32>) -> i32 { +// CHECK-NEXT: %c0_i32 = arith.constant 0 : i32 +// CHECK-NEXT: %0 = affine.parallel (%arg1) = (0) to (100) reduce ("maxs") -> (i32) { +// CHECK-NEXT: %2 = affine.load %arg0[%arg1] : memref<101xi32> +// CHECK-NEXT: %3 = affine.load %arg0[%arg1 + 1] : memref<101xi32> +// CHECK-NEXT: %4 = arith.subi %3, %2 : i32 +// CHECK-NEXT: affine.yield %4 : i32 +// CHECK-NEXT: } +// CHECK-NEXT: %1 = arith.maxsi %0, %c0_i32 : i32 +// CHECK-NEXT: return %1 : i32 +// CHECK-NEXT: } diff --git a/test/lit_tests/affinecfg_pad_loop_longest_row.mlir b/test/lit_tests/affinecfg_pad_loop_longest_row.mlir new file mode 100644 index 0000000000..c562de7d72 --- /dev/null +++ b/test/lit_tests/affinecfg_pad_loop_longest_row.mlir @@ -0,0 +1,296 @@ +// RUN: enzymexlamlir-opt --affine-cfg --split-input-file %s | FileCheck %s + +// A CSR transpose: the inner loop runs over a row's entries, its bounds +// read from the offsets, so it is not affine. It runs to the longest row +// instead, under `j < hi`, and the rows become its lanes. +func.func @csr_transpose(%off: memref, %idx: memref, %x: memref, %y: memref, %n: index) { + %cst = arith.constant 0.0 : f64 + %c1_i32 = arith.constant 1 : i32 + affine.parallel (%i) = (0) to (symbol(%n)) { + %lo = affine.load %off[%i] : memref + %hi = affine.load %off[%i + 1] : memref + %acc = scf.for %j = %lo to %hi step %c1_i32 iter_args(%a = %cst) -> (f64) : i32 { + %ji = arith.index_cast %j : i32 to index + %c = memref.load %idx[%ji] : memref + %ci = arith.index_cast %c : i32 to index + %v = memref.load %x[%ci] : memref + %s = arith.addf %a, %v : f64 + scf.yield %s : f64 + } + affine.store %acc, %y[%i] : memref + } + return +} + +// CHECK: func.func @csr_transpose(%arg0: memref, %arg1: memref, %arg2: memref, %arg3: memref, %arg4: index) { +// CHECK-NEXT: %cst = arith.constant -0.000000e+00 : f64 +// CHECK-NEXT: %0 = affine.parallel (%arg5) = (0) to (symbol(%arg4)) reduce ("maxs") -> (i32) { +// CHECK-NEXT: %2 = affine.load %arg0[%arg5] : memref +// CHECK-NEXT: %3 = affine.load %arg0[%arg5 + 1] : memref +// CHECK-NEXT: %4 = arith.subi %3, %2 : i32 +// CHECK-NEXT: affine.yield %4 : i32 +// CHECK-NEXT: } +// CHECK-NEXT: %1 = arith.index_cast %0 : i32 to index +// CHECK-NEXT: affine.parallel (%arg5) = (0) to (symbol(%arg4)) { +// CHECK-NEXT: %2 = affine.load %arg0[%arg5] : memref +// CHECK-NEXT: %3 = affine.load %arg0[%arg5 + 1] : memref +// CHECK-NEXT: %4 = affine.parallel (%arg6) = (0) to (symbol(%1)) reduce ("addf") -> (f64) { +// CHECK-NEXT: %5 = arith.index_cast %arg6 : index to i32 +// CHECK-NEXT: %6 = arith.addi %2, %5 : i32 +// CHECK-NEXT: %7 = arith.cmpi slt, %6, %3 : i32 +// CHECK-NEXT: %8 = scf.if %7 -> (f64) { +// CHECK-NEXT: %9 = arith.index_cast %6 : i32 to index +// CHECK-NEXT: %10 = memref.load %arg1[%9] : memref +// CHECK-NEXT: %11 = arith.index_cast %10 : i32 to index +// CHECK-NEXT: %12 = memref.load %arg2[%11] : memref +// CHECK-NEXT: scf.yield %12 : f64 +// CHECK-NEXT: } else { +// CHECK-NEXT: scf.yield %cst : f64 +// CHECK-NEXT: } +// CHECK-NEXT: affine.yield %8 : f64 +// CHECK-NEXT: } +// CHECK-NEXT: affine.store %4, %arg3[%arg5] : memref +// CHECK-NEXT: } +// CHECK-NEXT: return +// CHECK-NEXT: } + +// ----- + + + +// The rotated form clang makes of it, with a vector dimension between the +// row loop and the entries, and the bounds offset by one under the guard: +// the loop runs to the longest row of all the rows, and the guard stays. +func.func @rotated(%off: memref, %idx: memref, %x: memref, %y: memref, %n: index, %vdim: index) { + %cst = arith.constant 0.0 : f64 + %c1_i32 = arith.constant 1 : i32 + affine.parallel (%i) = (0) to (symbol(%n)) { + %lo = affine.load %off[%i] : memref + %hi = affine.load %off[%i + 1] : memref + %runs = arith.cmpi slt, %lo, %hi : i32 + affine.for %c = 0 to %vdim { + %acc = scf.if %runs -> (f64) { + %lo1 = arith.addi %lo, %c1_i32 : i32 + %hi1 = arith.addi %hi, %c1_i32 : i32 + %r = scf.for %j = %lo1 to %hi1 step %c1_i32 iter_args(%a = %cst) -> (f64) : i32 { + %jm = arith.subi %j, %c1_i32 : i32 + %ji = arith.index_cast %jm : i32 to index + %col = memref.load %idx[%ji] : memref + %ci = arith.index_cast %col : i32 to index + %v = memref.load %x[%ci] : memref + %s = arith.addf %a, %v : f64 + scf.yield %s : f64 + } + scf.yield %r : f64 + } else { + scf.yield %cst : f64 + } + affine.store %acc, %y[%i + %c * symbol(%n)] : memref + } + } + return +} + +// CHECK: func.func @rotated(%arg0: memref, %arg1: memref, %arg2: memref, %arg3: memref, %arg4: index, %arg5: index) { +// CHECK-NEXT: %cst = arith.constant -0.000000e+00 : f64 +// CHECK-NEXT: %cst_0 = arith.constant 0.000000e+00 : f64 +// CHECK-NEXT: %c1_i32 = arith.constant 1 : i32 +// CHECK-NEXT: %0 = affine.parallel (%arg6) = (0) to (symbol(%arg4)) reduce ("maxs") -> (i32) { +// CHECK-NEXT: %2 = affine.load %arg0[%arg6] : memref +// CHECK-NEXT: %3 = affine.load %arg0[%arg6 + 1] : memref +// CHECK-NEXT: %4 = arith.addi %2, %c1_i32 : i32 +// CHECK-NEXT: %5 = arith.addi %3, %c1_i32 : i32 +// CHECK-NEXT: %6 = arith.subi %5, %4 : i32 +// CHECK-NEXT: affine.yield %6 : i32 +// CHECK-NEXT: } +// CHECK-NEXT: %1 = arith.index_cast %0 : i32 to index +// CHECK-NEXT: affine.parallel (%arg6) = (0) to (symbol(%arg4)) { +// CHECK-NEXT: %2 = affine.load %arg0[%arg6] : memref +// CHECK-NEXT: %3 = affine.load %arg0[%arg6 + 1] : memref +// CHECK-NEXT: %4 = arith.cmpi slt, %2, %3 : i32 +// CHECK-NEXT: affine.for %arg7 = 0 to %arg5 { +// CHECK-NEXT: %5 = scf.if %4 -> (f64) { +// CHECK-NEXT: %6 = arith.addi %3, %c1_i32 : i32 +// CHECK-NEXT: %7 = affine.parallel (%arg8) = (0) to (symbol(%1)) reduce ("addf") -> (f64) { +// CHECK-NEXT: %8 = arith.index_cast %arg8 : index to i32 +// CHECK-NEXT: %9 = arith.addi %8, %2 : i32 +// CHECK-NEXT: %10 = arith.addi %9, %c1_i32 : i32 +// CHECK-NEXT: %11 = arith.cmpi slt, %10, %6 : i32 +// CHECK-NEXT: %12 = scf.if %11 -> (f64) { +// CHECK-NEXT: %13 = arith.index_cast %9 : i32 to index +// CHECK-NEXT: %14 = memref.load %arg1[%13] : memref +// CHECK-NEXT: %15 = arith.index_cast %14 : i32 to index +// CHECK-NEXT: %16 = memref.load %arg2[%15] : memref +// CHECK-NEXT: scf.yield %16 : f64 +// CHECK-NEXT: } else { +// CHECK-NEXT: scf.yield %cst : f64 +// CHECK-NEXT: } +// CHECK-NEXT: affine.yield %12 : f64 +// CHECK-NEXT: } +// CHECK-NEXT: scf.yield %7 : f64 +// CHECK-NEXT: } else { +// CHECK-NEXT: scf.yield %cst_0 : f64 +// CHECK-NEXT: } +// CHECK-NEXT: affine.store %5, %arg3[%arg6 + %arg7 * symbol(%arg4)] : memref +// CHECK-NEXT: } +// CHECK-NEXT: } +// CHECK-NEXT: return +// CHECK-NEXT: } + +// ----- + + + +// Bounds varying with two loops: the reduction is over both. +func.func @two_rows(%off: memref, %x: memref, %y: memref, %n: index) { + %cst = arith.constant 0.0 : f64 + %c1 = arith.constant 1 : index + affine.parallel (%i) = (0) to (symbol(%n)) { + affine.for %c = 0 to 3 { + %lo = affine.load %off[%i * 3 + %c] : memref + %hi = affine.load %off[%i * 3 + %c + 1] : memref + %loi = arith.index_cast %lo : i32 to index + %hii = arith.index_cast %hi : i32 to index + %acc = scf.for %j = %loi to %hii step %c1 iter_args(%a = %cst) -> (f64) { + %v = memref.load %x[%j] : memref + %s = arith.addf %a, %v : f64 + scf.yield %s : f64 + } + affine.store %acc, %y[%i * 3 + %c] : memref + } + } + return +} + +// CHECK: func.func @two_rows(%arg0: memref, %arg1: memref, %arg2: memref, %arg3: index) { +// CHECK-NEXT: %cst = arith.constant -0.000000e+00 : f64 +// CHECK-NEXT: %0 = affine.parallel (%arg4) = (0) to (symbol(%arg3)) reduce ("maxs") -> (i64) { +// CHECK-NEXT: %2 = affine.parallel (%arg5) = (0) to (3) reduce ("maxs") -> (i64) { +// CHECK-NEXT: %3 = affine.load %arg0[%arg5 + %arg4 * 3] : memref +// CHECK-NEXT: %4 = affine.load %arg0[%arg5 + %arg4 * 3 + 1] : memref +// CHECK-NEXT: %5 = arith.index_cast %3 : i32 to index +// CHECK-NEXT: %6 = arith.index_cast %4 : i32 to index +// CHECK-NEXT: %7 = arith.subi %6, %5 : index +// CHECK-NEXT: %8 = arith.index_cast %7 : index to i64 +// CHECK-NEXT: affine.yield %8 : i64 +// CHECK-NEXT: } +// CHECK-NEXT: affine.yield %2 : i64 +// CHECK-NEXT: } +// CHECK-NEXT: %1 = arith.index_cast %0 : i64 to index +// CHECK-NEXT: affine.parallel (%arg4) = (0) to (symbol(%arg3)) { +// CHECK-NEXT: affine.for %arg5 = 0 to 3 { +// CHECK-NEXT: %2 = affine.load %arg0[%arg5 + %arg4 * 3] : memref +// CHECK-NEXT: %3 = affine.load %arg0[%arg5 + %arg4 * 3 + 1] : memref +// CHECK-NEXT: %4 = arith.index_cast %2 : i32 to index +// CHECK-NEXT: %5 = arith.index_cast %3 : i32 to index +// CHECK-NEXT: %6 = affine.parallel (%arg6) = (0) to (symbol(%1)) reduce ("addf") -> (f64) { +// CHECK-NEXT: %7 = arith.addi %4, %arg6 : index +// CHECK-NEXT: %8 = arith.cmpi slt, %7, %5 : index +// CHECK-NEXT: %9 = scf.if %8 -> (f64) { +// CHECK-NEXT: %10 = memref.load %arg1[%7] : memref +// CHECK-NEXT: scf.yield %10 : f64 +// CHECK-NEXT: } else { +// CHECK-NEXT: scf.yield %cst : f64 +// CHECK-NEXT: } +// CHECK-NEXT: affine.yield %9 : f64 +// CHECK-NEXT: } +// CHECK-NEXT: affine.store %6, %arg2[%arg5 + %arg4 * 3] : memref +// CHECK-NEXT: } +// CHECK-NEXT: } +// CHECK-NEXT: return +// CHECK-NEXT: } + +// ----- + + + +// A bound read under a condition does not run on every row: the loop +// stays. +func.func @guarded_read(%off: memref, %flags: memref, %x: memref, %y: memref, %n: index) { + %cst = arith.constant 0.0 : f64 + %c1_i32 = arith.constant 1 : i32 + %c0_i32 = arith.constant 0 : i32 + affine.parallel (%i) = (0) to (symbol(%n)) { + %f = affine.load %flags[%i] : memref + %lo = affine.load %off[%i] : memref + %hi = scf.if %f -> (i32) { + %h = affine.load %off[%i + 1] : memref + scf.yield %h : i32 + } else { + scf.yield %c0_i32 : i32 + } + %acc = scf.for %j = %lo to %hi step %c1_i32 iter_args(%a = %cst) -> (f64) : i32 { + %ji = arith.index_cast %j : i32 to index + %v = memref.load %x[%ji] : memref + %s = arith.addf %a, %v : f64 + scf.yield %s : f64 + } + affine.store %acc, %y[%i] : memref + } + return +} + +// CHECK: func.func @guarded_read(%arg0: memref, %arg1: memref, %arg2: memref, %arg3: memref, %arg4: index) { +// CHECK-NEXT: %cst = arith.constant 0.000000e+00 : f64 +// CHECK-NEXT: %c1_i32 = arith.constant 1 : i32 +// CHECK-NEXT: %c0_i32 = arith.constant 0 : i32 +// CHECK-NEXT: affine.parallel (%arg5) = (0) to (symbol(%arg4)) { +// CHECK-NEXT: %0 = affine.load %arg1[%arg5] : memref +// CHECK-NEXT: %1 = affine.load %arg0[%arg5] : memref +// CHECK-NEXT: %2 = scf.if %0 -> (i32) { +// CHECK-NEXT: %4 = affine.load %arg0[%arg5 + 1] : memref +// CHECK-NEXT: scf.yield %4 : i32 +// CHECK-NEXT: } else { +// CHECK-NEXT: scf.yield %c0_i32 : i32 +// CHECK-NEXT: } +// CHECK-NEXT: %3 = scf.for %arg6 = %1 to %2 step %c1_i32 iter_args(%arg7 = %cst) -> (f64) : i32 { +// CHECK-NEXT: %4 = arith.index_cast %arg6 : i32 to index +// CHECK-NEXT: %5 = memref.load %arg2[%4] : memref +// CHECK-NEXT: %6 = arith.addf %arg7, %5 : f64 +// CHECK-NEXT: scf.yield %6 : f64 +// CHECK-NEXT: } +// CHECK-NEXT: affine.store %3, %arg3[%arg5] : memref +// CHECK-NEXT: } +// CHECK-NEXT: return +// CHECK-NEXT: } + +// ----- + + + +// The offsets are written in the row loop: the loop stays. +func.func @written(%off: memref, %x: memref, %y: memref, %n: index) { + %cst = arith.constant 0.0 : f64 + %c1_i32 = arith.constant 1 : i32 + affine.parallel (%i) = (0) to (symbol(%n)) { + %lo = affine.load %off[%i] : memref + %hi = affine.load %off[%i + 1] : memref + %acc = scf.for %j = %lo to %hi step %c1_i32 iter_args(%a = %cst) -> (f64) : i32 { + %ji = arith.index_cast %j : i32 to index + %v = memref.load %x[%ji] : memref + %s = arith.addf %a, %v : f64 + scf.yield %s : f64 + } + affine.store %acc, %y[%i] : memref + affine.store %hi, %off[%i] : memref + } + return +} + +// CHECK: func.func @written(%arg0: memref, %arg1: memref, %arg2: memref, %arg3: index) { +// CHECK-NEXT: %cst = arith.constant 0.000000e+00 : f64 +// CHECK-NEXT: %c1_i32 = arith.constant 1 : i32 +// CHECK-NEXT: affine.parallel (%arg4) = (0) to (symbol(%arg3)) { +// CHECK-NEXT: %0 = affine.load %arg0[%arg4] : memref +// CHECK-NEXT: %1 = affine.load %arg0[%arg4 + 1] : memref +// CHECK-NEXT: %2 = scf.for %arg5 = %0 to %1 step %c1_i32 iter_args(%arg6 = %cst) -> (f64) : i32 { +// CHECK-NEXT: %3 = arith.index_cast %arg5 : i32 to index +// CHECK-NEXT: %4 = memref.load %arg1[%3] : memref +// CHECK-NEXT: %5 = arith.addf %arg6, %4 : f64 +// CHECK-NEXT: scf.yield %5 : f64 +// CHECK-NEXT: } +// CHECK-NEXT: affine.store %2, %arg2[%arg4] : memref +// CHECK-NEXT: affine.store %1, %arg0[%arg4] : memref +// CHECK-NEXT: } +// CHECK-NEXT: return +// CHECK-NEXT: } diff --git a/test/lit_tests/raising/affine_to_stablehlo_int_minmax.mlir b/test/lit_tests/raising/affine_to_stablehlo_int_minmax.mlir new file mode 100644 index 0000000000..c52449cf9a --- /dev/null +++ b/test/lit_tests/raising/affine_to_stablehlo_int_minmax.mlir @@ -0,0 +1,85 @@ +// RUN: enzymexlamlir-opt %s --split-input-file --raise-affine-to-stablehlo --canonicalize | FileCheck %s + +// The longest row of a CSR structure, a signed max over its offsets. +func.func private @row_max(%off: memref<101xi32, 1>, %out: memref) { + %0 = affine.parallel (%i) = (0) to (100) reduce ("maxs") -> (i32) { + %lo = affine.load %off[%i] : memref<101xi32, 1> + %hi = affine.load %off[%i + 1] : memref<101xi32, 1> + %d = arith.subi %hi, %lo : i32 + affine.yield %d : i32 + } + affine.store %0, %out[] : memref + return +} + +// CHECK: func.func private @row_max_raised(%arg0: tensor<101xi32>, %arg1: tensor) -> (tensor<101xi32>, tensor) { +// CHECK-NEXT: %c = stablehlo.constant dense<-2147483648> : tensor +// CHECK-NEXT: %0 = stablehlo.slice %arg0 [0:100] : (tensor<101xi32>) -> tensor<100xi32> +// CHECK-NEXT: %1 = stablehlo.slice %arg0 [1:101] : (tensor<101xi32>) -> tensor<100xi32> +// CHECK-NEXT: %2 = arith.subi %1, %0 : tensor<100xi32> +// CHECK-NEXT: %3 = stablehlo.reduce(%2 init: %c) applies stablehlo.maximum across dimensions = [0] : (tensor<100xi32>, tensor) -> tensor +// CHECK-NEXT: %4 = stablehlo.dynamic_update_slice %arg1, %3 : (tensor, tensor) -> tensor +// CHECK-NEXT: return %arg0, %4 : tensor<101xi32>, tensor +// CHECK-NEXT: } + +// ----- + +func.func private @col_min(%in: memref<5x20xi64, 1>, %out: memref<20xi64, 1>) { + affine.parallel (%j) = (0) to (20) { + %0 = affine.parallel (%i) = (0) to (5) reduce ("mins") -> (i64) { + %v = affine.load %in[%i, %j] : memref<5x20xi64, 1> + affine.yield %v : i64 + } + affine.store %0, %out[%j] : memref<20xi64, 1> + } + return +} + +// CHECK: func.func private @col_min_raised(%arg0: tensor<5x20xi64>, %arg1: tensor<20xi64>) -> (tensor<5x20xi64>, tensor<20xi64>) { +// CHECK-NEXT: %c = stablehlo.constant dense<0> : tensor +// CHECK-NEXT: %c_0 = stablehlo.constant dense<9223372036854775807> : tensor +// CHECK-NEXT: %0 = stablehlo.reduce(%arg0 init: %c_0) applies stablehlo.minimum across dimensions = [0] : (tensor<5x20xi64>, tensor) -> tensor<20xi64> +// CHECK-NEXT: %1 = stablehlo.dynamic_update_slice %arg1, %0, %c : (tensor<20xi64>, tensor<20xi64>, tensor) -> tensor<20xi64> +// CHECK-NEXT: return %arg0, %1 : tensor<5x20xi64>, tensor<20xi64> +// CHECK-NEXT: } + +// ----- + +func.func private @col_minmax_unsigned(%in: memref<5x20xi32, 1>, %out: memref<2x20xi32, 1>) { + affine.parallel (%j) = (0) to (20) { + %0 = affine.parallel (%i) = (0) to (5) reduce ("maxu") -> (i32) { + %v = affine.load %in[%i, %j] : memref<5x20xi32, 1> + affine.yield %v : i32 + } + %1 = affine.parallel (%i) = (0) to (5) reduce ("minu") -> (i32) { + %v = affine.load %in[%i, %j] : memref<5x20xi32, 1> + affine.yield %v : i32 + } + affine.store %0, %out[0, %j] : memref<2x20xi32, 1> + affine.store %1, %out[1, %j] : memref<2x20xi32, 1> + } + return +} + +// CHECK: func.func private @col_minmax_unsigned_raised(%arg0: tensor<5x20xi32>, %arg1: tensor<2x20xi32>) -> (tensor<5x20xi32>, tensor<2x20xi32>) { +// CHECK-NEXT: %c = stablehlo.constant dense<1> : tensor +// CHECK-NEXT: %c_0 = stablehlo.constant dense<0> : tensor +// CHECK-NEXT: %c_1 = stablehlo.constant dense<-1> : tensor +// CHECK-NEXT: %c_2 = stablehlo.constant dense<0> : tensor +// CHECK-NEXT: %0 = stablehlo.reduce(%arg0 init: %c_2) across dimensions = [0] : (tensor<5x20xi32>, tensor) -> tensor<20xi32> +// CHECK-NEXT: reducer(%arg2: tensor, %arg3: tensor) { +// CHECK-NEXT: %7 = arith.maxui %arg2, %arg3 : tensor +// CHECK-NEXT: stablehlo.return %7 : tensor +// CHECK-NEXT: } +// CHECK-NEXT: %1 = stablehlo.reshape %arg0 : (tensor<5x20xi32>) -> tensor<5x20xi32> +// CHECK-NEXT: %2 = stablehlo.reduce(%1 init: %c_1) across dimensions = [0] : (tensor<5x20xi32>, tensor) -> tensor<20xi32> +// CHECK-NEXT: reducer(%arg2: tensor, %arg3: tensor) { +// CHECK-NEXT: %7 = arith.minui %arg2, %arg3 : tensor +// CHECK-NEXT: stablehlo.return %7 : tensor +// CHECK-NEXT: } +// CHECK-NEXT: %3 = stablehlo.broadcast_in_dim %0, dims = [1] : (tensor<20xi32>) -> tensor<1x20xi32> +// CHECK-NEXT: %4 = stablehlo.dynamic_update_slice %arg1, %3, %c_0, %c_0 : (tensor<2x20xi32>, tensor<1x20xi32>, tensor, tensor) -> tensor<2x20xi32> +// CHECK-NEXT: %5 = stablehlo.broadcast_in_dim %2, dims = [1] : (tensor<20xi32>) -> tensor<1x20xi32> +// CHECK-NEXT: %6 = stablehlo.dynamic_update_slice %4, %5, %c, %c_0 : (tensor<2x20xi32>, tensor<1x20xi32>, tensor, tensor) -> tensor<2x20xi32> +// CHECK-NEXT: return %arg0, %6 : tensor<5x20xi32>, tensor<2x20xi32> +// CHECK-NEXT: } diff --git a/workspace.bzl b/workspace.bzl index 5e2cf85d14..5286c37ecf 100644 --- a/workspace.bzl +++ b/workspace.bzl @@ -105,7 +105,7 @@ echo " llvm::Error evalPrintOp(PrintOp& op, InterpreterValue operand) {" >> thir # place copies and spills above the exec restore of an if/else join, # miscompiling kernels (llvm/llvm-project#222368). Drop once XLA's LLVM # includes it. - sed -i.bak0 "s/llvm:generated.patch\\\",/llvm:generated.patch\\\", \\\"\\/\\/:patches\\/llvm_amdgpu_bb_prolog.patch\\\", \\\"\\/\\/:patches\\/llvm_mlir_import_inrange_width.patch\\\", \\\"\\/\\/:patches\\/llvm_orc_unw_revert.patch\\\",/g" third_party/llvm/workspace.bzl + sed -i.bak0 "s/llvm:generated.patch\\\",/llvm:generated.patch\\\", \\\"\\/\\/:patches\\/llvm_amdgpu_bb_prolog.patch\\\", \\\"\\/\\/:patches\\/llvm_mlir_import_inrange_width.patch\\\", \\\"\\/\\/:patches\\/llvm_orc_unw_revert.patch\\\", \\\"\\/\\/:patches\\/llvm_affine_parallel_signless_minmax.patch\\\",/g" third_party/llvm/workspace.bzl """, """ sed -i.bak0 "s/tf_http_archive/http_archive/g" third_party/llvm/workspace.bzl