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
53 changes: 52 additions & 1 deletion src/enzyme_ad/jax/Passes/LLVMToAffineAccess.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -846,6 +846,56 @@ struct AffineIfDeadResults : public OpRewritePattern<affine::AffineIfOp> {
}
};

// An access through a view of the null pointer can only execute as
// undefined behavior, so it is dynamically dead: an optional buffer a kernel
// receives as null is always guarded by a flag. Fold loads to a zero of the
// element type and drop stores; the views and the captured null then die,
// instead of blocking the kernel on an untypable operand.
struct NullAccessFold : public OpRewritePattern<enzymexla::Pointer2MemrefOp> {
using OpRewritePattern::OpRewritePattern;

LogicalResult matchAndRewrite(enzymexla::Pointer2MemrefOp p2m,
PatternRewriter &rewriter) const override {
Value src = p2m.getSource();
while (auto asc = src.getDefiningOp<LLVM::AddrSpaceCastOp>())
src = asc.getArg();
if (!src.getDefiningOp<LLVM::ZeroOp>())
return failure();
bool changed = false;
for (Operation *user : llvm::make_early_inc_range(p2m->getUsers())) {
auto zeroFor = [&](Operation *ld, Type t) {
rewriter.setInsertionPoint(ld);
// getZeroAttr has no answer for pointers; those zero as llvm null.
if (t.isIntOrFloat()) {
rewriter.replaceOpWithNewOp<arith::ConstantOp>(
ld, t, rewriter.getZeroAttr(t));
return true;
}
if (isa<LLVM::LLVMPointerType>(t)) {
rewriter.replaceOpWithNewOp<LLVM::ZeroOp>(ld, t);
return true;
}
return false;
};
if (auto ld = dyn_cast<affine::AffineLoadOp>(user)) {
changed |= zeroFor(ld, ld.getType());
} else if (auto ld = dyn_cast<memref::LoadOp>(user)) {
changed |= zeroFor(ld, ld.getType());
} else if (isa<affine::AffineStoreOp, memref::StoreOp>(user) &&
user->getOperand(1) == p2m.getResult() &&
user->getOperand(0) != p2m.getResult()) {
rewriter.eraseOp(user);
changed = true;
}
}
if (p2m->use_empty()) {
rewriter.eraseOp(p2m);
changed = true;
}
return success(changed);
}
};

struct LoadSelect : public OpRewritePattern<affine::AffineLoadOp> {
using OpRewritePattern::OpRewritePattern;

Expand Down Expand Up @@ -2179,7 +2229,8 @@ convertLLVMToAffineAccess(Operation *op,
SimplifyDeadAlloc<memref::AllocaOp>, SimplifyDeadAlloc<memref::AllocOp>,
SimplifyDeadAlloc<LLVM::AllocaOp>,
SimplifyDeadAlloc<gpu::AllocOp, true>, Pointer2MemrefSelect, LoadSelect,
AffineIfDeadResults, SimpleMem2Reg<memref::AllocaOp>>(context);
NullAccessFold, AffineIfDeadResults, SimpleMem2Reg<memref::AllocaOp>>(
context);
GreedyRewriteConfig config;
config.setRegionSimplificationLevel(GreedySimplifyRegionLevel::Normal);
config.enableFolding();
Expand Down
43 changes: 43 additions & 0 deletions test/lit_tests/null_access_fold.mlir
Original file line number Diff line number Diff line change
@@ -0,0 +1,43 @@
// RUN: enzymexlamlir-opt %s --llvm-to-affine-access | FileCheck %s

// An optional buffer arriving as null: its accesses sit behind a runtime
// flag, so they can only execute as undefined behavior. Loads fold to zero,
// stores drop, and the null views disappear.
module {
func.func @optional(%out: memref<?xf64>, %flag: i1, %v: f64) {
%null = llvm.mlir.zero : !llvm.ptr
%fview = "enzymexla.pointer2memref"(%null) : (!llvm.ptr) -> memref<?xf64>
%iview = "enzymexla.pointer2memref"(%null) : (!llvm.ptr) -> memref<?xi32>
scf.if %flag {
%m = affine.load %fview[3] : memref<?xf64>
%i = affine.load %iview[1] : memref<?xi32>
%fi = arith.sitofp %i : i32 to f64
%s = arith.addf %m, %fi : f64
affine.store %s, %out[0] : memref<?xf64>
affine.store %v, %fview[2] : memref<?xf64>
}
return
}
}

// CHECK-LABEL: func.func @optional(
// CHECK-SAME: %[[OUT:.+]]: memref<?xf64>, %[[FLAG:.+]]: i1, %[[V:.+]]: f64
// CHECK-DAG: %[[FZ:.+]] = arith.constant 0.000000e+00 : f64
// CHECK-NOT: llvm.mlir.zero
// CHECK-NOT: pointer2memref
// CHECK: scf.if %[[FLAG]] {
// CHECK-NEXT: affine.store %[[FZ]], %[[OUT]][0] : memref<?xf64>
// CHECK-NEXT: }
// CHECK-NEXT: return

// A pointer loaded from a null view has no arith zero attribute; it folds to
// llvm null instead.
// CHECK-LABEL: func.func @optionalptr(
// CHECK: llvm.mlir.zero
// CHECK-NOT: arith.constant {{.*}} !llvm.ptr
func.func @optionalptr(%flag: i1) -> !llvm.ptr {
%null = llvm.mlir.zero : !llvm.ptr
%view = "enzymexla.pointer2memref"(%null) : (!llvm.ptr) -> memref<?x!llvm.ptr>
%p = affine.load %view[3] : memref<?x!llvm.ptr>
return %p : !llvm.ptr
}
Loading