Skip to content
Open
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
27 changes: 27 additions & 0 deletions patches/llvm_affine_parallel_signless_minmax.patch
Original file line number Diff line number Diff line change
@@ -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<FloatType>(resultType);
case arith::AtomicRMWKind::maxs: {
auto intType = dyn_cast<IntegerType>(resultType);
- return intType && intType.isSigned();
+ return intType && !intType.isUnsigned();
}
case arith::AtomicRMWKind::mins: {
auto intType = dyn_cast<IntegerType>(resultType);
- return intType && intType.isSigned();
+ return intType && !intType.isUnsigned();
}
case arith::AtomicRMWKind::maxu: {
auto intType = dyn_cast<IntegerType>(resultType);
- return intType && intType.isUnsigned();
+ return intType && !intType.isSigned();
}
case arith::AtomicRMWKind::minu: {
auto intType = dyn_cast<IntegerType>(resultType);
- return intType && intType.isUnsigned();
+ return intType && !intType.isSigned();
}
case arith::AtomicRMWKind::ori:
case arith::AtomicRMWKind::andi:
265 changes: 262 additions & 3 deletions src/enzyme_ad/jax/Passes/AffineCFG.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -3428,11 +3428,12 @@ struct ForOpRaising : public OpRewritePattern<scf::ForOp> {

// 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<AtomicRMWKind> reductionKind(Operation *op) {
auto intType = dyn_cast<IntegerType>(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<AtomicRMWKind> {
if (ok)
Expand Down Expand Up @@ -7331,6 +7332,263 @@ struct AffineForCopyCarry : public OpRewritePattern<AffineForOp> {
}
};

// The buffer behind a memref: the pointer a pointer2memref views, or the
// memref itself.
static Value bufferOf(Value memref) {
while (auto p2m = memref.getDefiningOp<enzymexla::Pointer2MemrefOp>())
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<OpTrait::HasRecursiveMemoryEffects>())
return;
auto iface = dyn_cast<MemoryEffectOpInterface>(op);
if (!iface) {
written = !isMemoryEffectFree(op);
return;
}
SmallVector<MemoryEffects::EffectInstance> effects;
iface.getEffects(effects);
for (auto &effect : effects) {
if (isa<MemoryEffects::Read, MemoryEffects::Allocate>(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<AffineForOp>(loop))
return forOp.getBody()->getArguments().take_front(1);
return cast<AffineParallelOp>(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<Operation *> &rows,
SetVector<Operation *> &chain) {
SmallVector<Value> todo{lo, hi};
DenseSet<Value> seen;
SetVector<Operation *> rowSet;
SmallVector<Value> loads;
while (!todo.empty()) {
Value v = todo.pop_back_val();
if (!seen.insert(v).second)
continue;
if (auto arg = dyn_cast<BlockArgument>(v)) {
Operation *owner = arg.getOwner()->getParentOp();
if (!owner->isProperAncestor(loop))
return false;
if (isa<AffineForOp, AffineParallelOp>(owner) &&
llvm::is_contained(loopIVs(owner), v)) {
rowSet.insert(owner);
continue;
}
if (isa<LoopLikeOpInterface>(owner))
return false;
continue;
}
Operation *op = v.getDefiningOp();
if (!scope->isAncestor(op->getParentRegion()))
continue;
if (isa<AffineLoadOp, memref::LoadOp>(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<Value> operands;
if (auto forOp = dyn_cast<AffineForOp>(row)) {
llvm::append_range(operands, forOp.getLowerBoundOperands());
llvm::append_range(operands, forOp.getUpperBoundOperands());
} else {
auto par = cast<AffineParallelOp>(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<BlockArgument>(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<scf::ForOp> {
using OpRewritePattern<scf::ForOp>::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<Operation *> rows;
SetVector<Operation *> 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<OpTrait::AffineScope>())
return failure();

Location loc = loop.getLoc();
Type ivType = loop.getInductionVar().getType();
Type redType = isa<IndexType>(ivType) ? rewriter.getI64Type() : ivType;

// The reduction: the rows again, yielding the row's length.
rewriter.setInsertionPoint(outer);
IRMapping map;
SmallVector<AffineParallelOp> reductions;
for (Operation *row : rows) {
SmallVector<AffineMap> lbMaps, ubMaps;
SmallVector<Value> lbArgs, ubArgs;
SmallVector<int64_t> steps;
if (auto forOp = dyn_cast<AffineForOp>(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<AffineParallelOp>(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<WalkOrder::PreOrder>([&](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<Value> 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<scf::ForOp> {
using OpRewritePattern<scf::ForOp>::OpRewritePattern;
Expand Down Expand Up @@ -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<PadLoopToLongestRow>(context, 0);
rpl.add</*SimplfyIntegerCastMath, */ CanonicalizeAffineApply, ForOpRaising,
ParallelOpRaising, CanonicalizeIndexCast<IndexCastOp>,
CanonicalizeIndexCast<IndexCastUIOp>, AffineIfYieldMovementPattern,
Expand Down
37 changes: 37 additions & 0 deletions src/enzyme_ad/jax/Passes/AffineToStableHLORaising.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<Attribute>(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<Attribute>(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<Attribute>(IntegerAttr::get(
ET, APInt::getAllOnes(ET.getIntOrFloatBitWidth())))));
innerRedName = "arith.minui";
break;
case arith::AtomicRMWKind::ori:
inits[0] = stablehlo::ConstantOp::create(
builder,
Expand Down
Loading
Loading